CoolFace
Apppublic

ImagingDataCommons/SegVolOnIDC

sourceHugging Facemitupdated 2y agoView on Hugging Face
2likes
app.py438 linesDownload Raw Back to root
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()