CoolFace
Apppublic

BAAI/SegVol

sourceHugging Facemitupdated 3y agoView on Hugging Face
8likes
app.py339 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, run12import nibabel as nib13import tempfile14 15print('script run')16 17#############################################18# init session_state19if 'option' not in  st.session_state:20    st.session_state.option = None21if 'text_prompt' not in st.session_state:22    st.session_state.text_prompt = None23 24if 'reset_demo_case' not in st.session_state:25    st.session_state.reset_demo_case = False26 27if 'preds_3D' not in st.session_state:28    st.session_state.preds_3D = None29    st.session_state.preds_3D_ori = None30 31if 'data_item' not in st.session_state:32    st.session_state.data_item = None33 34if 'points' not in st.session_state:35    st.session_state.points = []36 37if 'use_text_prompt' not in st.session_state:38    st.session_state.use_text_prompt = False39 40if 'use_point_prompt' not in st.session_state:41    st.session_state.use_point_prompt = False42 43if 'use_box_prompt' not in st.session_state:44    st.session_state.use_box_prompt = False45 46if 'rectangle_3Dbox' not in st.session_state:47    st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]48 49if 'irregular_box' not in st.session_state:50    st.session_state.irregular_box = False51 52if 'running' not in st.session_state:53    st.session_state.running = False54 55if 'transparency' not in st.session_state:56    st.session_state.transparency = 0.2557 58case_list = [59    'model/asset/FLARE22_Tr_0002_0000.nii.gz',60    'model/asset/FLARE22_Tr_0005_0000.nii.gz',61    'model/asset/FLARE22_Tr_0034_0000.nii.gz',62    'model/asset/FLARE22_Tr_0045_0000.nii.gz'63]64 65#############################################66 67#############################################68# reset functions69def clear_prompts():70    st.session_state.points = []71    st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]72 73def reset_demo_case():74    st.session_state.data_item = None75    st.session_state.reset_demo_case = True76    clear_prompts()77 78def clear_file():79    st.session_state.option = None80    process_ct_gt.clear()81    reset_demo_case()82    clear_prompts()83 84#############################################85 86st.image(Image.open('model/asset/overview back.png'), use_column_width=True)87 88github_col, arxive_col = st.columns(2)89 90with github_col:91    st.write('GitHub repo:https://github.com/BAAI-DCAI/SegVol')92 93with arxive_col:94    st.write('Paper:https://arxiv.org/abs/2311.13385')95 96 97# modify demo case here98demo_type = st.radio(99        "Demo case source",100        ["Select", "Upload"],101        on_change=clear_file102    )103 104if demo_type=="Select":105    uploaded_file = st.selectbox(106        "Select a demo case",107        case_list,108        index=None,109        placeholder="Select a demo case...",110        on_change=reset_demo_case111    )112else:113    uploaded_file = st.file_uploader("Upload demo case(nii.gz)", type='nii.gz', on_change=reset_demo_case)114 115st.session_state.option = uploaded_file116 117if  st.session_state.option is not None and \118    st.session_state.reset_demo_case or (st.session_state.data_item is None and st.session_state.option is not None):119 120    st.session_state.data_item = process_ct_gt(st.session_state.option)121    st.session_state.reset_demo_case = False122    st.session_state.preds_3D = None123    st.session_state.preds_3D_ori = None124 125prompt_col1, prompt_col2 = st.columns(2)126 127with prompt_col1:128    st.session_state.use_text_prompt = st.toggle('Sematic prompt')129    text_prompt_type = st.radio(130        "Sematic prompt type",131        ["Predefined", "Custom"],132        disabled=(not st.session_state.use_text_prompt)133    )134    if text_prompt_type == "Predefined":135        pre_text = st.selectbox(136            "Predefined anatomical category:",137            ['liver', 'right kidney', 'spleen', 'pancreas', 'aorta', 'inferior vena cava', 'right adrenal gland', 'left adrenal gland', 'gallbladder', 'esophagus', 'stomach', 'duodenum', 'left kidney'],138            index=None,139            disabled=(not st.session_state.use_text_prompt)140        )141    else:142        pre_text = st.text_input('Enter an Anatomical word or phrase:', None, max_chars=20,143                                                     disabled=(not st.session_state.use_text_prompt))144    if pre_text is None or len(pre_text) > 0:145        st.session_state.text_prompt = pre_text146    else:147        st.session_state.text_prompt = None148 149 150with prompt_col2:151    spatial_prompt_on = st.toggle('Spatial prompt', on_change=clear_prompts)152    spatial_prompt = st.radio(153        "Spatial prompt type",154        ["Point prompt", "Box prompt"],155        on_change=clear_prompts,156        disabled=(not spatial_prompt_on))157    st.session_state.enforce_zoom = st.checkbox('Enforce zoom-out-zoom-in')158 159if spatial_prompt == "Point prompt":160    st.session_state.use_point_prompt = True161    st.session_state.use_box_prompt = False162elif spatial_prompt == "Box prompt":163    st.session_state.use_box_prompt = True164    st.session_state.use_point_prompt = False165else:166    st.session_state.use_point_prompt = False167    st.session_state.use_box_prompt = False168 169if not spatial_prompt_on:170    st.session_state.use_point_prompt = False171    st.session_state.use_box_prompt = False172 173if not st.session_state.use_text_prompt:174    st.session_state.text_prompt = None175 176if st.session_state.option is None:177    st.write('please select demo case first')178else:179    image_3D = st.session_state.data_item['z_image'][0].numpy()180    col_control1, col_control2 = st.columns(2)181 182    with col_control1:183        selected_index_z = st.slider('X-Y view', 0, image_3D.shape[0] - 1, 162, key='xy', disabled=st.session_state.running)184 185    with col_control2:186        selected_index_y = st.slider('X-Z view', 0, image_3D.shape[1] - 1, 162, key='xz', disabled=st.session_state.running)187        if st.session_state.use_box_prompt:188            top, bottom = st.select_slider(189                'Top and bottom of box',190                options=range(0, 325),191                value=(0, 324), 192                disabled=st.session_state.running193            )194            st.session_state.rectangle_3Dbox[0] = top195            st.session_state.rectangle_3Dbox[3] = bottom196    col_image1, col_image2 = st.columns(2)197 198    if st.session_state.preds_3D is not None:199        st.session_state.transparency = st.slider('Mask opacity', 0.0, 1.0, 0.25, disabled=st.session_state.running)200 201    with col_image1:202        203        image_z_array = image_3D[selected_index_z]204 205        preds_z_array = None206        if st.session_state.preds_3D is not None:207            preds_z_array = st.session_state.preds_3D[selected_index_z]208            209        image_z = make_fig(image_z_array, preds_z_array, st.session_state.points, selected_index_z, 'xy')210        211        212        if st.session_state.use_point_prompt:213            value_xy = streamlit_image_coordinates(image_z, width=325)214            215            if value_xy is not None:216                point_ax_xy = (selected_index_z, value_xy['y'], value_xy['x'])217                if len(st.session_state.points) >= 3:218                    st.warning('Max point num is 3', icon="⚠️")219                elif point_ax_xy not in st.session_state.points:220                    st.session_state.points.append(point_ax_xy)221                    print('point_ax_xy add rerun')222                    st.rerun()223        elif st.session_state.use_box_prompt:224            canvas_result_xy = st_canvas(225                fill_color="rgba(255, 165, 0, 0.3)",  # Fixed fill color with some opacity226                stroke_width=3,227                stroke_color='#2909F1',228                background_image=image_z,229                update_streamlit=True,230                height=325,231                width=325,232                drawing_mode='transform',233                point_display_radius=0,234                key="canvas_xy",235                initial_drawing=initial_rectangle,236                display_toolbar=True237            )238            try:239                print(canvas_result_xy.json_data['objects'][0]['angle'])240                if canvas_result_xy.json_data['objects'][0]['angle'] != 0:241                    st.warning('Rotating is undefined behavior', icon="⚠️")242                    st.session_state.irregular_box = True243                else:244                    st.session_state.irregular_box = False245                reflect_json_data_to_3D_box(canvas_result_xy.json_data, view='xy')246            except:247                print('exception')248                pass249        else:250            st.image(image_z, use_column_width=False)251 252    with col_image2:253        image_y_array = image_3D[:, selected_index_y, :]254        255        preds_y_array = None256        if st.session_state.preds_3D is not None:257            preds_y_array = st.session_state.preds_3D[:, selected_index_y, :]258        259        image_y = make_fig(image_y_array, preds_y_array, st.session_state.points, selected_index_y, 'xz')260        261        if st.session_state.use_point_prompt:262            value_yz = streamlit_image_coordinates(image_y, width=325)263            264            if value_yz is not None:265                point_ax_xz = (value_yz['y'], selected_index_y, value_yz['x'])266                if len(st.session_state.points) >= 3:267                    st.warning('Max point num is 3', icon="⚠️")268                elif point_ax_xz not in st.session_state.points:269                    st.session_state.points.append(point_ax_xz)270                    print('point_ax_xz add rerun')271                    st.rerun()272        elif st.session_state.use_box_prompt:273            if st.session_state.rectangle_3Dbox[1] <= selected_index_y and selected_index_y <= st.session_state.rectangle_3Dbox[4]:274                draw = ImageDraw.Draw(image_y)275                #rectangle xz view (upper-left and lower-right)276                rectangle_coords = [(st.session_state.rectangle_3Dbox[2], st.session_state.rectangle_3Dbox[0]),277                                    (st.session_state.rectangle_3Dbox[5], st.session_state.rectangle_3Dbox[3])]278                # Draw the rectangle on the image279                draw.rectangle(rectangle_coords, outline='#2909F1', width=3)280            st.image(image_y, use_column_width=False)281        else:282            st.image(image_y, use_column_width=False)283 284 285col1, col2, col3 = st.columns(3)286 287with col1:288    if st.button("Clear", use_container_width=True,289                 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))):290        clear_prompts()291        st.session_state.preds_3D = None292        st.session_state.preds_3D_ori = None293        st.rerun()294 295with col2:296    img_nii = None297    if st.session_state.preds_3D_ori is not None and st.session_state.data_item is not None:298        meta_dict = st.session_state.data_item['meta']299        foreground_start_coord = st.session_state.data_item['foreground_start_coord']300        foreground_end_coord = st.session_state.data_item['foreground_end_coord']301        original_shape = st.session_state.data_item['ori_shape']302        pred_array = st.session_state.preds_3D_ori303        original_array = np.zeros(original_shape)304        original_array[foreground_start_coord[0]:foreground_end_coord[0], 305                    foreground_start_coord[1]:foreground_end_coord[1], 306                    foreground_start_coord[2]:foreground_end_coord[2]] = pred_array307 308        original_array = original_array.transpose(2, 1, 0)309        img_nii = nib.Nifti1Image(original_array, affine=meta_dict['affine'])310 311        with tempfile.NamedTemporaryFile(suffix=".nii.gz") as tmpfile:312            nib.save(img_nii, tmpfile.name)313            with open(tmpfile.name, "rb") as f:314                bytes_data = f.read()315                st.download_button(316                    label="Download result(.nii.gz)",317                    data=bytes_data,318                    file_name="segvol_preds.nii.gz",319                    mime="application/octet-stream",320                    disabled=img_nii is None321                )322 323with col3:324    run_button_name = 'Run'if not st.session_state.running else 'Running'325    if st.button(run_button_name, type="primary", use_container_width=True,326            disabled=(327                st.session_state.data_item is None or328                (st.session_state.text_prompt is None and len(st.session_state.points) == 0 and st.session_state.use_box_prompt is False) or 329                st.session_state.irregular_box or 330                st.session_state.running331                )):332        st.session_state.running = True333        st.rerun()334 335if st.session_state.running:336    st.session_state.running = False337    with st.status("Running...", expanded=False) as status:338        run()339    st.rerun()