CoolFace
Apppublic

yuxindu/SegVol

sourceHugging Facemitupdated 3y agoView on Hugging Face
4likes
app.py308 linesDownload Raw Back to root
1import streamlit as st2from streamlit_drawable_canvas import st_canvas3from streamlit_image_coordinates import streamlit_image_coordinates4 5 6from model.data_process.demo_data_process import process_ct_gt7import numpy as np8import matplotlib.pyplot as plt9from PIL import Image, ImageDraw10import monai.transforms as transforms11from utils import show_points, make_fig, reflect_points_into_model, initial_rectangle, reflect_json_data_to_3D_box, reflect_box_into_model, run12 13print('script run')14 15#############################################16# init session_state17if 'option' not in  st.session_state:18    st.session_state.option = None19if 'text_prompt' not in st.session_state:20    st.session_state.text_prompt = None21 22if 'reset_demo_case' not in st.session_state:23    st.session_state.reset_demo_case = False24 25if 'preds_3D' not in st.session_state:26    st.session_state.preds_3D = None27 28if 'data_item' not in st.session_state:29    st.session_state.data_item = None30 31if 'points' not in st.session_state:32    st.session_state.points = []33 34if 'use_text_prompt' not in st.session_state:35    st.session_state.use_text_prompt = False36 37if 'use_point_prompt' not in st.session_state:38    st.session_state.use_point_prompt = False39 40if 'use_box_prompt' not in st.session_state:41    st.session_state.use_box_prompt = False42 43if 'rectangle_3Dbox' not in st.session_state:44    st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]45 46if 'irregular_box' not in st.session_state:47    st.session_state.irregular_box = False48 49if 'running' not in st.session_state:50    st.session_state.running = False51 52if 'transparency' not in st.session_state:53    st.session_state.transparency = 0.2554 55case_list = [56    'model/asset/FLARE22_Tr_0002_0000.nii.gz',57    'model/asset/FLARE22_Tr_0005_0000.nii.gz',58    'model/asset/FLARE22_Tr_0034_0000.nii.gz',59    'model/asset/FLARE22_Tr_0045_0000.nii.gz'60]61 62#############################################63 64#############################################65# reset functions66def clear_prompts():67    st.session_state.points = []68    st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]69 70def reset_demo_case():71    st.session_state.data_item = None72    st.session_state.reset_demo_case = True73    clear_prompts()74 75def clear_file():76    st.session_state.option = None77    process_ct_gt.clear()78    reset_demo_case()79    clear_prompts()80 81#############################################82 83st.image(Image.open('model/asset/overview back.png'), use_column_width=True)84 85github_col, arxive_col = st.columns(2)86 87with github_col:88    st.write('GitHub repo:https://github.com/BAAI-DCAI/SegVol')89 90with arxive_col:91    st.write('Paper:https://arxiv.org/abs/2311.13385')92 93 94# modify demo case here95demo_type = st.radio(96        "Demo case source",97        ["Select", "Upload"],98        on_change=clear_file99    )100 101if demo_type=="Select":102    uploaded_file = st.selectbox(103        "Select a demo case",104        case_list,105        index=None,106        placeholder="Select a demo case...",107        on_change=reset_demo_case108    )109else:110    uploaded_file = st.file_uploader("Upload demo case(nii.gz)", type='nii.gz', on_change=reset_demo_case)111 112st.session_state.option = uploaded_file113 114if  st.session_state.option is not None and \115    st.session_state.reset_demo_case or (st.session_state.data_item is None and st.session_state.option is not None):116 117    st.session_state.data_item = process_ct_gt(st.session_state.option)118    st.session_state.reset_demo_case = False119    st.session_state.preds_3D = None120 121prompt_col1, prompt_col2 = st.columns(2)122 123with prompt_col1:124    st.session_state.use_text_prompt = st.toggle('Sematic prompt')125    text_prompt_type = st.radio(126        "Sematic prompt type",127        ["Predefined", "Custom"],128        disabled=(not st.session_state.use_text_prompt)129    )130    if text_prompt_type == "Predefined":131        pre_text = st.selectbox(132            "Predefined anatomical category:",133            ['liver', 'right kidney', 'spleen', 'pancreas', 'aorta', 'inferior vena cava', 'right adrenal gland', 'left adrenal gland', 'gallbladder', 'esophagus', 'stomach', 'duodenum', 'left kidney'],134            index=None,135            disabled=(not st.session_state.use_text_prompt)136        )137    else:138        pre_text = st.text_input('Enter an Anatomical word or phrase:', None, max_chars=20,139                                                     disabled=(not st.session_state.use_text_prompt))140    if pre_text is None or len(pre_text) > 0:141        st.session_state.text_prompt = pre_text142    else:143        st.session_state.text_prompt = None144 145 146with prompt_col2:147    spatial_prompt_on = st.toggle('Spatial prompt', on_change=clear_prompts)148    spatial_prompt = st.radio(149        "Spatial prompt type",150        ["Point prompt", "Box prompt"],151        on_change=clear_prompts,152        disabled=(not spatial_prompt_on))153 154if spatial_prompt == "Point prompt":155    st.session_state.use_point_prompt = True156    st.session_state.use_box_prompt = False157elif spatial_prompt == "Box prompt":158    st.session_state.use_box_prompt = True159    st.session_state.use_point_prompt = False160else:161    st.session_state.use_point_prompt = False162    st.session_state.use_box_prompt = False163 164if not spatial_prompt_on:165    st.session_state.use_point_prompt = False166    st.session_state.use_box_prompt = False167 168if not st.session_state.use_text_prompt:169    st.session_state.text_prompt = None170 171if st.session_state.option is None:172    st.write('please select demo case first')173else:174    image_3D = st.session_state.data_item['z_image'][0].numpy()175    col_control1, col_control2 = st.columns(2)176 177    with col_control1:178        selected_index_z = st.slider('X-Y view', 0, image_3D.shape[0] - 1, 162, key='xy', disabled=st.session_state.running)179 180    with col_control2:181        selected_index_y = st.slider('X-Z view', 0, image_3D.shape[1] - 1, 162, key='xz', disabled=st.session_state.running)182        if st.session_state.use_box_prompt:183            top, bottom = st.select_slider(184                'Top and bottom of box',185                options=range(0, 325),186                value=(0, 324), 187                disabled=st.session_state.running188            )189            st.session_state.rectangle_3Dbox[0] = top190            st.session_state.rectangle_3Dbox[3] = bottom191    col_image1, col_image2 = st.columns(2)192 193    if st.session_state.preds_3D is not None:194        st.session_state.transparency = st.slider('Mask opacity', 0.0, 1.0, 0.25, disabled=st.session_state.running)195 196    with col_image1:197        198        image_z_array = image_3D[selected_index_z]199 200        preds_z_array = None201        if st.session_state.preds_3D is not None:202            preds_z_array = st.session_state.preds_3D[selected_index_z]203            204        image_z = make_fig(image_z_array, preds_z_array, st.session_state.points, selected_index_z, 'xy')205        206        207        if st.session_state.use_point_prompt:208            value_xy = streamlit_image_coordinates(image_z, width=325)209            210            if value_xy is not None:211                point_ax_xy = (selected_index_z, value_xy['y'], value_xy['x'])212                if len(st.session_state.points) >= 3:213                    st.warning('Max point num is 3', icon="⚠️")214                elif point_ax_xy not in st.session_state.points:215                    st.session_state.points.append(point_ax_xy)216                    print('point_ax_xy add rerun')217                    st.rerun()218        elif st.session_state.use_box_prompt:219            canvas_result_xy = st_canvas(220                fill_color="rgba(255, 165, 0, 0.3)",  # Fixed fill color with some opacity221                stroke_width=3,222                stroke_color='#2909F1',223                background_image=image_z,224                update_streamlit=True,225                height=325,226                width=325,227                drawing_mode='transform',228                point_display_radius=0,229                key="canvas_xy",230                initial_drawing=initial_rectangle,231                display_toolbar=True232            )233            try:234                print(canvas_result_xy.json_data['objects'][0]['angle'])235                if canvas_result_xy.json_data['objects'][0]['angle'] != 0:236                    st.warning('Rotating is undefined behavior', icon="⚠️")237                    st.session_state.irregular_box = True238                else:239                    st.session_state.irregular_box = False240                reflect_json_data_to_3D_box(canvas_result_xy.json_data, view='xy')241            except:242                print('exception')243                pass244        else:245            st.image(image_z, use_column_width=False)246 247    with col_image2:248        image_y_array = image_3D[:, selected_index_y, :]249        250        preds_y_array = None251        if st.session_state.preds_3D is not None:252            preds_y_array = st.session_state.preds_3D[:, selected_index_y, :]253        254        image_y = make_fig(image_y_array, preds_y_array, st.session_state.points, selected_index_y, 'xz')255        256        if st.session_state.use_point_prompt:257            value_yz = streamlit_image_coordinates(image_y, width=325)258            259            if value_yz is not None:260                point_ax_xz = (value_yz['y'], selected_index_y, value_yz['x'])261                if len(st.session_state.points) >= 3:262                    st.warning('Max point num is 3', icon="⚠️")263                elif point_ax_xz not in st.session_state.points:264                    st.session_state.points.append(point_ax_xz)265                    print('point_ax_xz add rerun')266                    st.rerun()267        elif st.session_state.use_box_prompt:268            if st.session_state.rectangle_3Dbox[1] <= selected_index_y and selected_index_y <= st.session_state.rectangle_3Dbox[4]:269                draw = ImageDraw.Draw(image_y)270                #rectangle xz view (upper-left and lower-right)271                rectangle_coords = [(st.session_state.rectangle_3Dbox[2], st.session_state.rectangle_3Dbox[0]),272                                    (st.session_state.rectangle_3Dbox[5], st.session_state.rectangle_3Dbox[3])]273                # Draw the rectangle on the image274                draw.rectangle(rectangle_coords, outline='#2909F1', width=3)275            st.image(image_y, use_column_width=False)276        else:277            st.image(image_y, use_column_width=False)278 279 280col1, col2, col3 = st.columns(3)281 282with col1:283    if st.button("Clear", use_container_width=True,284                 disabled=(st.session_state.option is None or (len(st.session_state.points)==0 and not st.session_state.use_box_prompt and st.session_state.preds_3D is None))):285        clear_prompts()286        st.session_state.preds_3D = None287        st.rerun()288 289with col3:290    run_button_name = 'Run'if not st.session_state.running else 'Running'291    if st.button(run_button_name, type="primary", use_container_width=True,292            disabled=(293                st.session_state.data_item is None or294                (st.session_state.text_prompt is None and len(st.session_state.points) == 0 and st.session_state.use_box_prompt is False) or 295                st.session_state.irregular_box or 296                st.session_state.running297                )):298        st.session_state.running = True299        st.rerun()300 301# if len(st.session_state.points) > 0:302#     st.write(st.session_state.points)303 304if st.session_state.running:305    st.session_state.running = False306    with st.status("Running...", expanded=False) as status:307        run()308    st.rerun()