{"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":"# Interactive Viewing of Scans and Segmentation Data\n\nThis notebook allows the user to slide through each plane in the series with the option to view the segmentations provided by TotalSegmentator overlaid. The user must copy and run the notebook to view the scans. Selecting a series will load every scan into data and, if available, the segmentation data - hence the user should expect some number of seconds to pass before the first scan is displayed.\n\nNote: If the image freezes, then reloading the series will likely solve the problem.","metadata":{}},{"cell_type":"code","source":"# Load libraries.\nimport os\nimport glob\n\nimport joblib\nimport pandas as pd\nimport numpy as np\nimport nibabel as nib\nimport ipywidgets as widgets\nimport matplotlib.pyplot as plt\nimport pydicom","metadata":{"execution":{"iopub.status.busy":"2023-08-27T11:13:49.154556Z","iopub.execute_input":"2023-08-27T11:13:49.156025Z","iopub.status.idle":"2023-08-27T11:13:49.802006Z","shell.execute_reply.started":"2023-08-27T11:13:49.155639Z","shell.execute_reply":"2023-08-27T11:13:49.800778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the train and test meta data.\ntrain_meta = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train_series_meta.csv')\ntest_meta = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/test_series_meta.csv')","metadata":{"execution":{"iopub.status.busy":"2023-08-27T11:13:49.804298Z","iopub.execute_input":"2023-08-27T11:13:49.804691Z","iopub.status.idle":"2023-08-27T11:13:49.836328Z","shell.execute_reply.started":"2023-08-27T11:13:49.80464Z","shell.execute_reply":"2023-08-27T11:13:49.83516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dicom_tag_columns = [\n    'BitsAllocated',\n    'BitsStored',\n    'Columns',\n    'ImageOrientationPatient',\n    'ImagePositionPatient',\n    'InstanceNumber',\n    'PatientID',\n    'PatientPosition',\n    'PixelRepresentation',\n    'PixelSpacing',\n    'RescaleIntercept',\n    'RescaleSlope',\n    'Rows',\n    'SeriesNumber',\n    'SeriesInstanceUID',\n    'SliceThickness',\n    'path',\n    'WindowCenter',\n    'WindowWidth'\n]\nROOT_DIR = '/kaggle/input/rsna-2023-abdominal-trauma-detection/'\nVECTOR_DTYPE = np.float32\n\ndef process_dicom_tags(df, copy=True):\n    df_out = df.copy() if copy else df\n    df_out['SeriesID'] = df_out.path.str.split('/').str.get(2).astype(int)\n    # Cast strings as arrays.\n    for c in ['ImageOrientationPatient', 'ImagePositionPatient', 'PixelSpacing']:\n        df_out[c] = df_out[c].str.strip('[]').str.encode('utf-8').apply(lambda x: np.fromstring(x, sep=',', dtype=VECTOR_DTYPE))\n    # Make `path` absolute path.\n    df_out['path'] = df_out.path.apply(lambda x: os.path.join(ROOT_DIR, x))\n    df_out.sort_values(['SeriesID', 'InstanceNumber'], inplace=True)\n    return df_out","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-08-27T11:13:49.838413Z","iopub.execute_input":"2023-08-27T11:13:49.838938Z","iopub.status.idle":"2023-08-27T11:13:49.849974Z","shell.execute_reply.started":"2023-08-27T11:13:49.83889Z","shell.execute_reply":"2023-08-27T11:13:49.8486Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load DICOM tags and process them.\ntrain_dicom_tags = pd.read_parquet('/kaggle/input/rsna-2023-abdominal-trauma-detection/train_dicom_tags.parquet', columns=dicom_tag_columns)\ntest_dicom_tags = pd.read_parquet('/kaggle/input/rsna-2023-abdominal-trauma-detection/test_dicom_tags.parquet', columns=dicom_tag_columns)\ntrain_dicom_tags = process_dicom_tags(train_dicom_tags, copy=False)\ntest_dicom_tags = process_dicom_tags(test_dicom_tags, copy=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-27T11:13:49.853059Z","iopub.execute_input":"2023-08-27T11:13:49.853452Z","iopub.status.idle":"2023-08-27T11:14:28.511867Z","shell.execute_reply.started":"2023-08-27T11:13:49.85342Z","shell.execute_reply":"2023-08-27T11:14:28.510351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEGMENTATION_CODES = {\n    1: 'liver',\n    2: 'spleen',\n    3: 'kidney_left',\n    4: 'kidney_right',\n    5: 'bowel',\n}\n\ndef calibrate_img(img, m, c):\n    # Transform into HU.\n    img *= m\n    img += c\n    #img[img < hu_min] = hu_min # Clip pixels outside the FOV. \n    return img\n\n\ndef get_patient_id(series_id, use_train_set):\n    df = train_meta if use_train_set else test_meta\n    return int(df[df.series_id == series_id].iloc[0].patient_id)\n\n\ndef get_scan_files(patient_id, series_id, use_train_set):\n    dataset_dir = 'train_images' if use_train_set else 'test_images'\n    fnames = glob.glob(f'/kaggle/input/rsna-2023-abdominal-trauma-detection/{dataset_dir}/{patient_id}/{series_id}/*.dcm')\n    return list(sorted(fnames, key=lambda f: int(os.path.splitext(os.path.basename(f))[0])))\n\n\ndef load_series(series_id, use_train_set):\n    df = train_dicom_tags if use_train_set else test_dicom_tags\n    sub_df = df[df.SeriesID == series_id]\n\n    @joblib.delayed\n    def load_dcm(path, bit_shift):\n        dcm = pydicom.dcmread(path)\n        pixel_array = dcm.pixel_array\n        if bit_shift is not None:\n            dtype = pixel_array.dtype \n            pixel_array = (pixel_array << bit_shift).astype(dtype) >> bit_shift\n        return pixel_array.astype(np.float32)\n    \n    m, c, iop, (dx, dy), dz, pr, ba, bs = sub_df.iloc[0][\n        ['RescaleSlope', 'RescaleIntercept', 'ImageOrientationPatient', 'PixelSpacing', 'SliceThickness', 'PixelRepresentation', 'BitsAllocated', 'BitsStored']\n    ]\n    bit_shift = None\n    if pr == 1:\n        bit_shift = ba - bs\n\n    img = joblib.Parallel(n_jobs=-1)(load_dcm(path, bit_shift) for path in sub_df.path.values)\n    img = np.dstack(img)\n    \n    # Make sure images are ordered correctly (superior is +z).\n    img = calibrate_img(img, m, c)\n    imaging_axis = np.cross(iop[:3], iop[3:])\n    distance_projection = np.dot(np.vstack(sub_df.ImagePositionPatient.values), imaging_axis)\n    img = img[:, :, np.argsort(distance_projection)]\n    img = img.transpose((1, 0, 2))\n    spacing = (float(dx), float(dy), float(dz))\n    return img, spacing\n\n\ndef get_segmentation_paths():\n    paths = glob.glob(f'/kaggle/input/rsna-2023-abdominal-trauma-detection/segmentations/*.nii')\n    return list(sorted(paths, key=lambda f: int(os.path.splitext(os.path.basename(f))[0])))\n\n\ndef get_series_ids_with_segmentations():\n    return [int(os.path.splitext(os.path.basename(f))[0]) for f in get_segmentation_paths()]\n    \n\ndef get_segmentation_path(series_id):\n    return f'/kaggle/input/rsna-2023-abdominal-trauma-detection/segmentations/{series_id}.nii'\n\n\ndef load_segmentation(x):\n    if not isinstance(x, str):\n        x = get_segmentation_path(int(x))\n    out = nib.load(x)\n    return nib.as_closest_canonical(out) # Ensures output is RAS, DICOM is LPS.\n\n\ndef get_segmentation_image(seg):\n    if not isinstance(seg, nib.nifti1.Nifti1Image):\n        seg = load_segmentation(seg)\n    return seg.get_fdata().astype(int)[::-1, ::-1, :]","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-08-27T11:14:28.51377Z","iopub.execute_input":"2023-08-27T11:14:28.514641Z","iopub.status.idle":"2023-08-27T11:14:28.537188Z","shell.execute_reply.started":"2023-08-27T11:14:28.514597Z","shell.execute_reply":"2023-08-27T11:14:28.535771Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.patches import Patch\nfrom matplotlib.lines import Line2D\n\nSEGMENTATION_CMAP = plt.get_cmap('rainbow')\n\n\ndef capture_output(f):\n    def wrapper(self, *args, **kwargs):\n        with self.output_widget:\n            return f(self, *args, **kwargs)\n    return wrapper\n\n\n# Define widgets.\nclass CTWidget(widgets.VBox):\n    DATASET_WIDGET_OPT_TRAIN = 0\n    DATASET_WIDGET_OPT_TEST = 1\n    DATASET_WIDGET_OPTIONS = {\n        'Train': DATASET_WIDGET_OPT_TRAIN,\n        'Test': DATASET_WIDGET_OPT_TEST\n    }\n    VIEW_PLANE_WIDGET_OPT_XY = 0\n    VIEW_PLANE_WIDGET_OPT_YZ = 1\n    VIEW_PLANE_WIDGET_OPT_XZ = 2\n    VIEW_PLANE_WIDGET_OPTIONS = {\n        'Transverse (XY)': VIEW_PLANE_WIDGET_OPT_XY,\n        'Sagittal (YZ)': VIEW_PLANE_WIDGET_OPT_YZ,\n        'Coronal (XZ)': VIEW_PLANE_WIDGET_OPT_XZ,\n    }\n    VIEW_PLANE_CONSTANT_INDEXES = {\n        VIEW_PLANE_WIDGET_OPT_XY: 2,\n        VIEW_PLANE_WIDGET_OPT_XZ: 1,\n        VIEW_PLANE_WIDGET_OPT_YZ: 0,\n    }\n    SERIES_IDS_WITH_SEGMENTATIONS = get_series_ids_with_segmentations()\n    \n    SEGMENTATION_PALETTE = np.array([[np.nan] * 4] + [SEGMENTATION_CMAP(x) for x in np.linspace(0, 1, 6)])\n\n    def __init__(self):\n        super().__init__()\n        self.children = self.init_widgets()\n        self.current_series_data = None\n        self.current_series_spacing = None\n        self.current_segmentation_data = None\n        self.fig, self.ax = plt.subplots()\n        self.ax_im = None \n        self.ax_im_seg = None\n        plt.close(self.fig)\n\n    def init_widgets(self):\n        self.dataset_widget = widgets.Dropdown(\n            options=list(self.DATASET_WIDGET_OPTIONS),\n            value=None,\n            description='Dataset',\n            disabled=False\n        )\n        self.dataset_widget.observe(self.handle_dataset_update, names='value')\n        self.filter_series_id_widget = widgets.Checkbox(\n            value=False,\n            description='Filter series ID without segmentations',\n            disabled=False,\n            indent=False\n        )\n        self.filter_series_id_widget.observe(self.handle_filter_series_update, names='value')\n        self.series_id_widget = widgets.Dropdown(\n            options=[''],\n            description='Series ID',\n            disabled=False\n        )\n        self.series_id_widget.observe(self.handle_series_id_update, names='value')\n        self.view_plane_widget = widgets.RadioButtons(\n            options=list(self.VIEW_PLANE_WIDGET_OPTIONS),\n            layout={'width': 'max-content'}, # If the items' names are long\n            description='View Plane:',\n            disabled=False\n        )\n        self.view_plane_widget.observe(self.handle_view_plane_update, names='value')\n        self.slice_ix_widget = widgets.IntSlider(\n            value=0,\n            min=0,\n            max=1,\n            step=1,\n            description='Slice Index',\n            disabled=False,\n            continuous_update=True,\n            orientation='horizontal',\n            readout=True,\n            readout_format='d'\n        )\n        self.slice_ix_widget.observe(self.handle_slice_ix_update, names='value')\n        self.segmentation_enable_widget = widgets.Checkbox(\n            value=False,\n            description='Segmentations not available.',\n            disabled=True,\n            indent=False\n        )\n        self.segmentation_enable_widget.observe(self.handle_segmentation_enable_update, names='value')\n        self.window_range_widget = widgets.FloatRangeSlider(\n            value=[-1024.0, 1024.0],\n            min=-1024.0,\n            max=1024.0,\n            step=0.1,\n            description='Window Range (HU):',\n            disabled=False,\n            continuous_update=True,\n            orientation='horizontal',\n            readout=True,\n            readout_format='.1f',\n            style={'description_width': 'initial'}\n        )\n        self.window_range_widget.observe(self.handle_window_range_update, names='value')\n        self.segmentation_alpha_widget = widgets.FloatSlider(\n            value=0.5,\n            min=0.0,\n            max=1.0,\n            step=0.01,\n            description='Segmentation Alpha:',\n            disabled=False,\n            continuous_update=True,\n            orientation='horizontal',\n            readout=True,\n            readout_format='.2f',\n            style={'description_width': 'initial'}\n        )\n        self.segmentation_alpha_widget.observe(self.handle_segmentation_alpha_update, names='value')\n        self.output_widget = widgets.Output()\n        return [\n            self.dataset_widget,\n            self.filter_series_id_widget,\n            self.series_id_widget,\n            self.view_plane_widget,\n            self.slice_ix_widget,\n            self.segmentation_enable_widget,\n            self.window_range_widget,\n            self.segmentation_alpha_widget,\n            self.output_widget,\n        ]\n    \n    @property\n    def use_train_set(self):\n        return self.DATASET_WIDGET_OPTIONS[self.dataset_widget.value] == self.DATASET_WIDGET_OPT_TRAIN\n\n    @capture_output\n    def update_series_ids_options(self):\n        if self.dataset_widget.value is None:\n            return\n        if self.use_train_set and self.filter_series_id_widget.value:\n            opts = self.SERIES_IDS_WITH_SEGMENTATIONS\n        elif self.use_train_set and not self.filter_series_id_widget.value:\n            opts = sorted(train_meta.series_id.values)\n        elif not self.use_train_set and self.filter_series_id_widget.value:\n            opts = ['']\n        else:\n            opts = sorted(test_meta.series_id.values)\n        self.series_id_widget.options = opts\n        self.current_series_data = None\n        self.current_series_spacing = None\n        self.current_segmentation_data = None\n\n    @capture_output\n    def handle_segmentation_alpha_update(self, change):\n        if self.ax_im_seg is not None:\n            self.ax_im_seg.set_alpha(change.new)\n            self.fig.canvas.draw_idle()\n\n    @capture_output\n    def update_scan_ix_options(self):\n        lower = 0\n        if self.current_series_data is None:\n            upper = 1\n        else:\n            view_plane = self.VIEW_PLANE_WIDGET_OPTIONS[self.view_plane_widget.value]\n            ix = self.VIEW_PLANE_CONSTANT_INDEXES[view_plane]\n            upper = self.current_series_data.shape[ix] - 1\n        self.slice_ix_widget.min = lower\n        self.slice_ix_widget.max = upper\n        self.slice_ix_widget.value = 0\n\n    @capture_output\n    def handle_dataset_update(self, change):\n        # Update series IDs, clear value.\n        self.update_series_ids_options()\n        # Reset slice.\n        self.slice_ix_widget.value = 0\n\n    @capture_output\n    def handle_filter_series_update(self, change):\n        self.update_series_ids_options()\n    \n    @capture_output\n    def handle_series_id_update(self, change):\n        series_id = change.new\n        if series_id in self.SERIES_IDS_WITH_SEGMENTATIONS:\n            self.segmentation_enable_widget.disabled = False\n            self.segmentation_enable_widget.description = 'Show Segmentations'\n        else:\n            self.segmentation_enable_widget.value = False\n            self.segmentation_enable_widget.disabled = True\n            self.segmentation_enable_widget.description = 'Segmentations not available.'\n\n        self.clear_image()\n        self.current_series_data, self.current_series_spacing = load_series(series_id, self.use_train_set)\n        self.update_scan_ix_options()\n        if series_id in self.SERIES_IDS_WITH_SEGMENTATIONS:\n            self.current_segmentation_data = get_segmentation_image(series_id)\n        else:\n            self.current_segmentation_data = None\n\n    @capture_output\n    def handle_view_plane_update(self, change):\n        self.update_scan_ix_options()\n        self.clear_image()\n    \n    @capture_output\n    def handle_slice_ix_update(self, change):\n        pass\n\n    @capture_output\n    def handle_segmentation_enable_update(self, change):\n        if not change.new and self.ax_im_seg is not None:\n            self.ax_im_seg.remove()\n            self.ax_im_seg = None\n            if self.ax.get_legend() is not None:\n                self.ax.get_legend().remove()\n            self.fig.canvas.draw_idle()\n\n    @capture_output\n    def handle_window_range_update(self, change):\n        pass\n\n    @capture_output\n    def get_series_indexing(self):\n        indexing = [slice(0, None)] * 3\n        view_plane = self.VIEW_PLANE_WIDGET_OPTIONS[self.view_plane_widget.value]\n        ix = self.VIEW_PLANE_CONSTANT_INDEXES[view_plane]\n        indexing[ix] = self.slice_ix_widget.value\n        return tuple(indexing)\n\n    @capture_output\n    def get_extent(self):\n        view_plane = self.VIEW_PLANE_WIDGET_OPTIONS[self.view_plane_widget.value]\n        constant_ix = self.VIEW_PLANE_CONSTANT_INDEXES[view_plane]\n        axes = [ax for ax in range(3) if ax != constant_ix]\n        shapes = [self.current_series_data.shape[ax] for ax in axes]\n        spacings = [self.current_series_spacing[ax] for ax in axes]\n        return (0, shapes[0] * spacings[0], 0, shapes[1] * spacings[1]) # (left, right, bottom, top),\n    \n    def clear_image(self):\n        if self.ax_im is not None:\n            self.ax_im.remove()\n            self.ax_im = None\n            self.fig.canvas.draw_idle()\n        if self.ax_im_seg is not None:\n            self.ax_im_seg.remove()\n            self.ax_im_seg = None\n            if self.ax.get_legend() is not None:\n                self.ax.get_legend().remove()\n            self.fig.canvas.draw_idle()\n\n        \n    def update_image(self, *args, **kwargs):\n        if self.current_series_data is None:\n            return\n        indexing = self.get_series_indexing()\n        series_img = self.current_series_data[indexing].T\n        series_img = np.clip(series_img, a_min=self.window_range_widget.value[0], a_max=self.window_range_widget.value[1])\n        if self.ax_im is None:\n            self.ax_im = self.ax.imshow(series_img, cmap=plt.cm.bone, extent=self.get_extent())\n        else:\n            self.ax_im.set_data(series_img)\n            self.ax_im.set_clim(series_img.min(), series_img.max())\n\n        if self.current_segmentation_data is None or (self.current_segmentation_data is not None and not self.segmentation_enable_widget.value):\n            display(self.fig)\n            return\n        seg_img = self.current_segmentation_data[indexing].T\n        seg_img_coloured = self.SEGMENTATION_PALETTE[seg_img]\n        if self.ax_im_seg is None:\n            self.ax_im_seg = self.ax.imshow(seg_img_coloured, alpha=self.segmentation_alpha_widget.value, extent=self.get_extent())\n        else:\n            self.ax_im_seg.set_data(seg_img_coloured)\n        mask_levels = [x for x in np.unique(seg_img) if x != 0]\n        legend_elements = [\n            Patch(facecolor=self.SEGMENTATION_PALETTE[x], edgecolor=self.SEGMENTATION_PALETTE[x], label=SEGMENTATION_CODES[x]) for x in mask_levels\n        ]\n        if self.ax.get_legend() is not None:\n            self.ax.get_legend().remove()\n        self.ax.legend(handles=legend_elements, loc='upper right')\n        display(self.fig)\n\n    def interact(self):\n        out = widgets.interactive(\n            self.update_image,\n            dataset=self.dataset_widget,\n            filter_series_id=self.filter_series_id_widget,\n            series_id=self.series_id_widget,\n            view_plane=self.view_plane_widget,\n            slice_ix=self.slice_ix_widget,\n            show_segmentations=self.segmentation_enable_widget,\n            segmentations_alpha=self.segmentation_alpha_widget, \n            window_range=self.window_range_widget,\n        )\n        out.children[-1].layout.height = '450px'\n        return display(out)\n    \nw = CTWidget()\nw.interact()","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-08-27T11:14:28.539598Z","iopub.execute_input":"2023-08-27T11:14:28.540151Z","iopub.status.idle":"2023-08-27T11:14:28.738793Z","shell.execute_reply.started":"2023-08-27T11:14:28.540105Z","shell.execute_reply":"2023-08-27T11:14:28.737648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}