ImagingDataCommons/SegVolOnIDC
2
1import streamlit as st2from streamlit_drawable_canvas import st_canvas3from streamlit_image_coordinates import streamlit_image_coordinates4from idc_index import index5import os6import glob7import shutil8import dcm2niix9import subprocess10import random11import base6412 13from model.data_process.demo_data_process import process_ct_gt14import numpy as np15import matplotlib.pyplot as plt16from PIL import Image, ImageDraw17import monai.transforms as transforms18from utils import show_points, make_fig, reflect_points_into_model, initial_rectangle, reflect_json_data_to_3D_box, reflect_box_into_model, run19import nibabel as nib20import tempfile21 22print('script run')23#further improvement24#decorator singletion or use cache data class25# https://docs.streamlit.io/develop/api-reference/caching-and-state/st.experimental_singleton26# https://docs.streamlit.io/develop/concepts/architecture/caching27def download_idc_data_serieUID(serieUID_lst, output_folder):28 #download IDC data cases29 client = index.IDCClient()30 #define serieUIDs to download31 #download series and convert to .nii.gz32 if os.path.exists(output_folder):33 shutil.rmtree(output_folder)34 os.makedirs(output_folder)35 for idx, serieUID_ddl in enumerate(serieUID_lst):36 sample_dcm_dir = os.path.join(output_folder, f"ddl_series{idx}_dcm")37 sample_nii_dir = os.path.join(output_folder, f"ddl_series{idx}_nii")38 for dir in [sample_dcm_dir, sample_nii_dir]:39 if os.path.exists(dir):40 shutil.rmtree(dir)41 os.makedirs(dir)42 client.download_from_selection(seriesInstanceUID=serieUID_ddl, downloadDir=sample_dcm_dir)43 subprocess.call(["dcm2niix", "-o", sample_nii_dir, "-z", "y",44 "-f", "IDC_%i", "-g", "y", sample_dcm_dir])45 return glob.glob(os.path.join(output_folder, "*nii/*.nii.gz"))46 47 48def get_random_sample_idc_from_bodypart(bodypart_selected):49 client = index.IDCClient()50 # body_parts = client.index[(client.index['Modality'].isin(['CT']))&(idc_client.index['instanceCount']> '100')]['BodyPartExamined'].unique()51 matching_series_list = client.index[client.index['Modality'].isin(["CT"]) \52 & (client.index['BodyPartExamined'] == bodypart_selected) & \53 (client.index['instanceCount']> '100')]['SeriesInstanceUID'].values54 # select random series from the list55 random_series_uid = random.choice(matching_series_list)56 random_series_viewer_url = client.get_viewer_URL(random_series_uid)57 return random_series_uid, random_series_viewer_url58 59def retrieve_idc_index_body_parts():60 idc_client = index.IDCClient()61 body_parts = idc_client.index[(idc_client.index['Modality'].isin(['CT']))&(idc_client.index['instanceCount']< '150')]['BodyPartExamined'].unique()62 return body_parts63 64#############################################65st.session_state.option = None66 67if 'idc_data' not in st.session_state:68 case_list = download_idc_data_serieUID(serieUID_lst=["1.3.6.1.4.1.14519.5.2.1.8421.4008.125612661111422710051062993644",69 "1.3.6.1.4.1.14519.5.2.1.3344.4008.552105302448832783460360105045",70 "1.3.6.1.4.1.14519.5.2.1.3344.4008.217290429362492484143666931850",71 "1.3.6.1.4.1.14519.5.2.1.3344.4008.315023636447426194723399171147",72 "1.3.6.1.4.1.14519.5.2.1.3344.4008.307374355712319704057189924161"],73 output_folder="model/asset/idc_samples")74 st.session_state.idc_data = True75else:76 case_list = glob.glob("model/asset/idc_samples/*nii/*.nii.gz")77if 'idc_serieUID_sample' not in st.session_state:78 st.session_state.idc_serieUID_sample = None79# init session_state80if 'option' not in st.session_state:81 st.session_state.option = None82if 'text_prompt' not in st.session_state:83 st.session_state.text_prompt = None84if 'reset_demo_case' not in st.session_state:85 st.session_state.reset_demo_case = False86 87if 'preds_3D' not in st.session_state:88 st.session_state.preds_3D = None89 st.session_state.preds_3D_ori = None90 91if 'data_item' not in st.session_state:92 st.session_state.data_item = None93 94if 'points' not in st.session_state:95 st.session_state.points = []96 97if 'use_text_prompt' not in st.session_state:98 st.session_state.use_text_prompt = False99 100if 'use_text_serieUID' not in st.session_state:101 st.session_state.use_text_serieUID = False102 103if 'use_point_prompt' not in st.session_state:104 st.session_state.use_point_prompt = False105 106if 'use_box_prompt' not in st.session_state:107 st.session_state.use_box_prompt = False108 109if 'rectangle_3Dbox' not in st.session_state:110 st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]111 112if 'irregular_box' not in st.session_state:113 st.session_state.irregular_box = False114 115if 'running' not in st.session_state:116 st.session_state.running = False117 118if 'transparency' not in st.session_state:119 st.session_state.transparency = 0.25120#############################################121 122#############################################123# reset functions124def clear_prompts():125 st.session_state.points = []126 st.session_state.rectangle_3Dbox = [0,0,0,0,0,0]127 128def reset_demo_case():129 st.session_state.data_item = None130 st.session_state.idc_serieUID_sample = None131 st.session_state.reset_demo_case = True132 st.session_state.idc_bodypart_selected = False133 clear_prompts()134 135def clear_file():136 st.session_state.option = None137 st.session_state.idc_serieUID_sample = None138 st.session_state.idc_bodypart_selected = False139 process_ct_gt.clear()140 reset_demo_case()141 clear_prompts()142 143#############################################144st.image("idc_intro_extended.jpg")145st.write("Below is an example on how to select a SeriesInstanceUID from Imaging Data Commons (IDC) to further use in this demo:")146st.image("https://github.com/ccosmin97/huggingface_idc_demos/raw/main/idc_serieUID_selection.gif")147st.write("Below is an overview of the SegVol method and authors acknowledgement.")148st.image(Image.open('model/asset/overview back.png'), use_column_width=True)149 150github_col, arxive_col = st.columns(2)151 152with github_col:153 st.write('SegVol GitHub repo:https://github.com/BAAI-DCAI/SegVol')154 155with arxive_col:156 st.write('SegVol Paper:https://arxiv.org/abs/2311.13385')157 158 159# modify demo case here160demo_type = st.radio(161 "Demo case source",162 ["Select an IDC demo case from tcga_lihc collection", 163 "Filter by DICOM SeriesInstanceUID", 164 "Random sampling based on BodyPartExamined"],165 on_change=clear_file166 )167 168if demo_type=="Select an IDC demo case from tcga_lihc collection":169 uploaded_file = st.selectbox(170 "Select a demo case",171 case_list,172 index=None,173 placeholder="Select a demo case...",174 on_change=reset_demo_case)175elif demo_type=="Filter by DICOM SeriesInstanceUID":176 with st.form("Filter by DICOM SeriesInstanceUID"):177 uploaded_serieUID = st.text_input("Enter a DICOM SeriesInstanceUID", value=None)178 submitted = st.form_submit_button("Submit", on_click=clear_prompts)179 if submitted:180 st.session_state.idc_serieUID_sample = download_idc_data_serieUID([str(uploaded_serieUID).strip()], "model/asset/idc_serieUID_sample")[0]181 # st.session_state.option = uploaded_file182 uploaded_file = st.session_state.idc_serieUID_sample183 else:184 uploaded_file = st.session_state.idc_serieUID_sample185else:#elif demo_type == "Random sampling based on BodyPartExamined":186 with st.form("Filter by DICOM BodyPartExamined Tag") as form_body_part:187 # body_part_list = retrieve_idc_index_body_parts()188 body_part_selected = st.selectbox(189 "Select a bodypart to randomly sample a CT scan from",190 ["ABDOMEN", "LUNG", "LIVER",191 "PELVIS"],192 index=None,193 placeholder="Select a bodypart to pick a SeriesInstanceUID from...")194 submitted = st.form_submit_button("Submit", on_click=reset_demo_case)195 #if st.session_state.reset_demo_case == True and body_part_selected is not None:# and st.session_state.idc_bodypart_selected == False and 196 if submitted: 197 serieUID, ohif_link = get_random_sample_idc_from_bodypart(body_part_selected)198 for i in range(0,5):199 if os.path.exists("model/asset/idc_serieUID_random_sample"):200 shutil.rmtree("model/asset/idc_serieUID_random_sample")201 st.session_state.idc_serieUID_sample = download_idc_data_serieUID([str(serieUID)], "model/asset/idc_serieUID_random_sample")[0]202 path_file = glob.glob(f"model/asset/idc_serieUID_random_sample/ddl_series0_nii/*.nii.gz")203 if path_file and len(path_file) == 1:204 break205 else:206 print("serieUID NOT FILLING BASIC REQs --> MORE THAN 1 NII FILE OR NO NII FILE")207 # st.write(f"SeriesInstanceUID randomly sampled from chosen BodyPartExamined : {random_series_uid}")208 # st.write(f"OHIF URL of selected sample : {random_series_viewer_url}")209 # st.session_state.idc_bodypart_selected = True210 uploaded_file = st.session_state.idc_serieUID_sample211 else:212 uploaded_file = st.session_state.idc_serieUID_sample213 214st.session_state.option = uploaded_file 215 216if st.session_state.option is not None and \217 st.session_state.reset_demo_case or (st.session_state.data_item is None and st.session_state.option is not None):218 219 st.session_state.data_item = process_ct_gt(st.session_state.option)220 st.session_state.reset_demo_case = False221 st.session_state.preds_3D = None222 st.session_state.preds_3D_ori = None223 224prompt_col1, prompt_col2 = st.columns(2)225 226with prompt_col1:227 st.session_state.use_text_prompt = st.toggle('Semantic prompt')228 text_prompt_type = st.radio(229 "Semantic prompt type",230 ["Predefined", "Custom"],231 disabled=(not st.session_state.use_text_prompt)232 )233 if text_prompt_type == "Predefined":234 pre_text = st.selectbox(235 "Predefined anatomical category:",236 ['liver', 'right kidney', 'spleen', 'pancreas', 'aorta', 'inferior vena cava', 'right adrenal gland', 'left adrenal gland', 'gallbladder', 'esophagus', 'stomach', 'duodenum', 'left kidney'],237 index=None,238 disabled=(not st.session_state.use_text_prompt)239 )240 else:241 pre_text = st.text_input('Enter an Anatomical word or phrase:', None, max_chars=20,242 disabled=(not st.session_state.use_text_prompt))243 if pre_text is None or len(pre_text) > 0:244 st.session_state.text_prompt = pre_text245 else:246 st.session_state.text_prompt = None247 248 249with prompt_col2:250 spatial_prompt_on = st.toggle('Spatial prompt', on_change=clear_prompts)251 spatial_prompt = st.radio(252 "Spatial prompt type",253 ["Point prompt", "Box prompt"],254 on_change=clear_prompts,255 disabled=(not spatial_prompt_on))256 st.session_state.enforce_zoom = st.checkbox('Enforce zoom-out-zoom-in')257 258if spatial_prompt == "Point prompt":259 st.session_state.use_point_prompt = True260 st.session_state.use_box_prompt = False261elif spatial_prompt == "Box prompt":262 st.session_state.use_box_prompt = True263 st.session_state.use_point_prompt = False264else:265 st.session_state.use_point_prompt = False266 st.session_state.use_box_prompt = False267 268if not spatial_prompt_on:269 st.session_state.use_point_prompt = False270 st.session_state.use_box_prompt = False271 272if not st.session_state.use_text_prompt:273 st.session_state.text_prompt = None274 275if st.session_state.option is None:276 st.write('please select demo case first')277else:278 image_3D = st.session_state.data_item['z_image'][0].numpy()279 col_control1, col_control2 = st.columns(2)280 281 with col_control1:282 selected_index_z = st.slider('X-Y view', 0, image_3D.shape[0] - 1, 162, key='xy', disabled=st.session_state.running)283 284 with col_control2:285 selected_index_y = st.slider('X-Z view', 0, image_3D.shape[1] - 1, 162, key='xz', disabled=st.session_state.running)286 if st.session_state.use_box_prompt:287 top, bottom = st.select_slider(288 'Top and bottom of box',289 options=range(0, 325),290 value=(0, 324),291 disabled=st.session_state.running292 )293 st.session_state.rectangle_3Dbox[0] = top294 st.session_state.rectangle_3Dbox[3] = bottom295 col_image1, col_image2 = st.columns(2)296 297 if st.session_state.preds_3D is not None:298 st.session_state.transparency = st.slider('Mask opacity', 0.0, 1.0, 0.25, disabled=st.session_state.running)299 300 with col_image1:301 302 image_z_array = image_3D[selected_index_z]303 304 preds_z_array = None305 if st.session_state.preds_3D is not None:306 preds_z_array = st.session_state.preds_3D[selected_index_z]307 308 image_z = make_fig(image_z_array, preds_z_array, st.session_state.points, selected_index_z, 'xy')309 310 311 if st.session_state.use_point_prompt:312 value_xy = streamlit_image_coordinates(image_z, width=325)313 314 if value_xy is not None:315 point_ax_xy = (selected_index_z, value_xy['y'], value_xy['x'])316 if len(st.session_state.points) >= 3:317 st.warning('Max point num is 3', icon="??")318 elif point_ax_xy not in st.session_state.points:319 st.session_state.points.append(point_ax_xy)320 print('point_ax_xy add rerun')321 st.rerun()322 elif st.session_state.use_box_prompt:323 canvas_result_xy = st_canvas(324 fill_color="rgba(255, 165, 0, 0.3)", # Fixed fill color with some opacity325 stroke_width=3,326 stroke_color='#2909F1',327 background_image=image_z,328 update_streamlit=True,329 height=325,330 width=325,331 drawing_mode='transform',332 point_display_radius=0,333 key="canvas_xy",334 initial_drawing=initial_rectangle,335 display_toolbar=True336 )337 try:338 print(canvas_result_xy.json_data['objects'][0]['angle'])339 if canvas_result_xy.json_data['objects'][0]['angle'] != 0:340 st.warning('Rotating is undefined behavior', icon="??")341 st.session_state.irregular_box = True342 else:343 st.session_state.irregular_box = False344 reflect_json_data_to_3D_box(canvas_result_xy.json_data, view='xy')345 except:346 print('exception')347 pass348 else:349 st.image(image_z, use_column_width=False)350 351 with col_image2:352 image_y_array = image_3D[:, selected_index_y, :]353 354 preds_y_array = None355 if st.session_state.preds_3D is not None:356 preds_y_array = st.session_state.preds_3D[:, selected_index_y, :]357 358 image_y = make_fig(image_y_array, preds_y_array, st.session_state.points, selected_index_y, 'xz')359 360 if st.session_state.use_point_prompt:361 value_yz = streamlit_image_coordinates(image_y, width=325)362 363 if value_yz is not None:364 point_ax_xz = (value_yz['y'], selected_index_y, value_yz['x'])365 if len(st.session_state.points) >= 3:366 st.warning('Max point num is 3', icon="??")367 elif point_ax_xz not in st.session_state.points:368 st.session_state.points.append(point_ax_xz)369 print('point_ax_xz add rerun')370 st.rerun()371 elif st.session_state.use_box_prompt:372 if st.session_state.rectangle_3Dbox[1] <= selected_index_y and selected_index_y <= st.session_state.rectangle_3Dbox[4]:373 draw = ImageDraw.Draw(image_y)374 #rectangle xz view (upper-left and lower-right)375 rectangle_coords = [(st.session_state.rectangle_3Dbox[2], st.session_state.rectangle_3Dbox[0]),376 (st.session_state.rectangle_3Dbox[5], st.session_state.rectangle_3Dbox[3])]377 # Draw the rectangle on the image378 draw.rectangle(rectangle_coords, outline='#2909F1', width=3)379 st.image(image_y, use_column_width=False)380 else:381 st.image(image_y, use_column_width=False)382 383 384col1, col2, col3 = st.columns(3)385 386with col1:387 if st.button("Clear", use_container_width=True,388 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))):389 clear_prompts()390 st.session_state.preds_3D = None391 st.session_state.preds_3D_ori = None392 st.rerun()393 394with col2:395 img_nii = None396 if st.session_state.preds_3D_ori is not None and st.session_state.data_item is not None:397 meta_dict = st.session_state.data_item['meta']398 foreground_start_coord = st.session_state.data_item['foreground_start_coord']399 foreground_end_coord = st.session_state.data_item['foreground_end_coord']400 original_shape = st.session_state.data_item['ori_shape']401 pred_array = st.session_state.preds_3D_ori402 original_array = np.zeros(original_shape)403 original_array[foreground_start_coord[0]:foreground_end_coord[0],404 foreground_start_coord[1]:foreground_end_coord[1],405 foreground_start_coord[2]:foreground_end_coord[2]] = pred_array406 407 original_array = original_array.transpose(2, 1, 0)408 img_nii = nib.Nifti1Image(original_array, affine=meta_dict['affine'])409 410 with tempfile.NamedTemporaryFile(suffix=".nii.gz") as tmpfile:411 nib.save(img_nii, tmpfile.name)412 with open(tmpfile.name, "rb") as f:413 bytes_data = f.read()414 st.download_button(415 label="Download result(.nii.gz)",416 data=bytes_data,417 file_name="segvol_preds.nii.gz",418 mime="application/octet-stream",419 disabled=img_nii is None420 )421 422with col3:423 run_button_name = 'Run'if not st.session_state.running else 'Running'424 if st.button(run_button_name, type="primary", use_container_width=True,425 disabled=(426 st.session_state.data_item is None or427 (st.session_state.text_prompt is None and len(st.session_state.points) == 0 and st.session_state.use_box_prompt is False) or428 st.session_state.irregular_box or429 st.session_state.running430 )):431 st.session_state.running = True432 st.rerun()433 434if st.session_state.running:435 st.session_state.running = False436 with st.status("Running...", expanded=False) as status:437 run()438 st.rerun()