{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Plotting  interactive 3D segmentation masks using `plotly` package","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nimport gc\n\nimport numpy as np\nimport pandas as pd\n\nimport plotly.offline as py\nimport plotly.graph_objs as go\nimport plotly.express as px\n\nimport scipy.ndimage\nimport nibabel as nib\n\nTRAIN_CSV_PATH = '../input/rsna-2022-cervical-spine-fracture-detection/train.csv'\nSEG_PATH = '../input/rsna-2022-cervical-spine-fracture-detection/segmentations/'\ntrain=pd.read_csv(TRAIN_CSV_PATH)\nfrac_cols = train.columns[-7:]","metadata":{"execution":{"iopub.status.busy":"2022-08-17T11:52:11.303001Z","iopub.execute_input":"2022-08-17T11:52:11.303415Z","iopub.status.idle":"2022-08-17T11:52:11.317187Z","shell.execute_reply.started":"2022-08-17T11:52:11.303379Z","shell.execute_reply":"2022-08-17T11:52:11.315983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_patients = [x[:-4] for x in os.listdir(SEG_PATH)]\ntrain['Fractures'] = train[frac_cols].sum(axis=1)\ntrain['Has_mask'] = train.StudyInstanceUID.isin(seg_patients).astype('int').map({0:'No', 1:'Yes'})\n\nfig = px.histogram(train, x='Fractures', color='Has_mask', log_y=True,\n                   title=\"Number of fractures by patient\", color_discrete_sequence=[\"gray\", \"crimson\"])\\\n        .update_yaxes(categoryorder='total descending', title='Number of patients (Log scale)')\\\n        .update_layout(bargap=0.1)   \n    \nfig.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-17T11:53:09.253997Z","iopub.execute_input":"2022-08-17T11:53:09.254426Z","iopub.status.idle":"2022-08-17T11:53:09.337175Z","shell.execute_reply.started":"2022-08-17T11:53:09.254389Z","shell.execute_reply":"2022-08-17T11:53:09.335903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transform 3D mask to coordinates","metadata":{}},{"cell_type":"code","source":"#number of segmentation masks to plot\nn_cases = 15\n\npatients = os.listdir(SEG_PATH)\npatients = np.random.choice(patients, n_cases, replace=False)\n\ndata = pd.DataFrame()\nfor p in patients:\n    img = nib.load(SEG_PATH + p).get_fdata()\n    \n    img = scipy.ndimage.zoom(img, 0.5).astype(np.int8)\n    x, y, z = img.nonzero()\n    values = [img[P[0], P[1], P[2]] for P in zip(x,y,z)]\n    \n    df = pd.DataFrame({'x': x,'y': y,'z': z,'val': values})\n    df['patient'] = p\n    \n    data = pd.concat([data, df])\n    \ndata","metadata":{"execution":{"iopub.status.busy":"2022-08-17T11:45:47.948148Z","iopub.execute_input":"2022-08-17T11:45:47.948559Z","iopub.status.idle":"2022-08-17T11:48:44.934585Z","shell.execute_reply.started":"2022-08-17T11:45:47.948523Z","shell.execute_reply":"2022-08-17T11:48:44.932998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Plot segmentation masks","metadata":{}},{"cell_type":"code","source":"fig_data = [] \nbuttons = []\n\ndef plot_scatters_3d(data):\n    \n    patients = data.patient.unique()\n    \n    for i, p in enumerate(patients):\n        df = data[data.patient == p]\n\n        scatter = go.Scatter3d(name=f'{p}', x=df.x,y=df.y,z=df.z, mode='markers', \n                               marker=dict(color = df.val, opacity=0.75, size = 2, colorscale='cividis'),\n                               visible = not bool(i))\n\n        fig_data.append(scatter)\n\n        visible = [False]*len(patients)\n        visible[i] = True\n        buttons.append(\n            dict(method='restyle', args=[{'visible': visible}], label=p)\n        )    \n\n    layout = go.Layout(width=800, height=800, title='Segmentation masks', updatemenus=list([\n        dict(x=-0.05 ,y=1, buttons=buttons, yanchor='top', xanchor=\"left\", showactive=True)\n        ]))\n\n    fig = dict(data=fig_data, layout=layout)\n    \n    return fig\n    \nscatter = plot_scatters_3d(data)\npy.iplot(scatter)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T11:48:44.937715Z","iopub.execute_input":"2022-08-17T11:48:44.938164Z","iopub.status.idle":"2022-08-17T11:48:52.758091Z","shell.execute_reply.started":"2022-08-17T11:48:44.938125Z","shell.execute_reply":"2022-08-17T11:48:52.755899Z"},"trusted":true},"execution_count":null,"outputs":[]}]}