{"metadata":{"environment":{"kernel":"conda-env-kaggle-rsna-2024-env-kaggle-rsna-2024-env","name":"workbench-notebooks.m119","type":"gcloud","uri":"us-docker.pkg.dev/deeplearning-platform-release/gcr.io/workbench-notebooks:m119"},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":8516623,"sourceType":"datasetVersion","datasetId":5084734},{"sourceId":179913984,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA 2024 Lumbar Spine Degenerative Classification - DICOM Viewer\n\n\nThis notebook allows the user to slide through each slice in each plane of a given series. To run the viewer, fork this notebook and click `Run >> Run all`.\n\nTo change the plotting backend, change `Config.plotting_backend` to one of `[\"matplotlib\", \"plotly\"]`.","metadata":{}},{"cell_type":"code","source":"class Config:\n    plotting_backend = \"matplotlib\"  # one of \"matplotlib\", \"plotly\"","metadata":{"execution":{"iopub.status.busy":"2024-06-30T20:57:36.793464Z","iopub.execute_input":"2024-06-30T20:57:36.793976Z","iopub.status.idle":"2024-06-30T20:57:36.832784Z","shell.execute_reply.started":"2024-06-30T20:57:36.793934Z","shell.execute_reply":"2024-06-30T20:57:36.831658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\ntry:\n    import fastparquet\nexcept ModuleNotFoundError:\n    os.system(\"pip install -q fastparquet\")\n    \nfrom rsna_24_utilities import FileSystem\n\nfs = FileSystem()\ntrain_dicom_tags_df = fs.dicom_tags.load_dicom_tags(train=True)\ntest_dicom_tags_df = fs.dicom_tags.load_dicom_tags(train=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-30T20:57:36.834971Z","iopub.execute_input":"2024-06-30T20:57:36.835763Z","iopub.status.idle":"2024-06-30T20:58:07.504225Z","shell.execute_reply.started":"2024-06-30T20:57:36.835723Z","shell.execute_reply":"2024-06-30T20:58:07.503155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Any\n\nfrom rsna_24_utilities import StrEnum\n\n\nclass ViewPlane(StrEnum):\n    XY = \"Transverse (XY)\"\n    YZ = \"Sagittal (YZ)\"\n    XZ = \"Coronal (XZ)\"\n\n    @classmethod\n    def _axes_mapping(cls) -> dict:\n        return {\n            cls.XY: (0, 1),\n            cls.YZ: (1, 2),\n            cls.XZ: (0, 2),\n        }\n\n    @classmethod\n    def _projected_axis_mapping(cls) -> dict:\n        return {\n            cls.XY: 2,\n            cls.YZ: 0,\n            cls.XZ: 1,\n        }\n\n    @property\n    def axes(self) -> tuple[int, int]:\n        return self._axes_mapping()[self]\n\n    @property\n    def projected_axis(self) -> int:\n        return self._projected_axis_mapping()[self]\n\n    def get_slice(self, ix):\n        slicing = [slice(None) for _ in range(3)]\n        slicing[self.projected_axis] = ix\n        return tuple(slicing)\n\n\nclass Dataset(StrEnum):\n    TRAIN = \"Train\"\n    TEST = \"Test\"\n\n    @property\n    def is_train(self) -> bool:\n        return self is type(self).TRAIN\n\n\nclass PlottingBackend(StrEnum):\n    MPL = \"matplotlib\"\n    PLOTLY = \"plotly\"","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-06-30T20:58:07.505745Z","iopub.execute_input":"2024-06-30T20:58:07.506286Z","iopub.status.idle":"2024-06-30T20:58:07.517247Z","shell.execute_reply.started":"2024-06-30T20:58:07.506253Z","shell.execute_reply":"2024-06-30T20:58:07.516093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ipywidgets as widgets\nimport matplotlib.pyplot as plt\nimport plotly.graph_objects as go\nfrom matplotlib import colormaps as cm\nfrom matplotlib.colors import Normalize\n\n\nclass ViewPlaneWidget(widgets.VBox):\n    def __init__(self, parent, col: int, view_plane: ViewPlane, plotting_backend: PlottingBackend):\n        super().__init__()\n        self.parent = parent\n        self.col = col\n        self.view_plane = view_plane\n        self.plotting_backend = plotting_backend\n\n        if plotting_backend == PlottingBackend.MPL:\n            self.fig, self.ax = plt.subplots()\n            plt.close()\n            self.im = None\n        elif plotting_backend == PlottingBackend.PLOTLY:\n            self.fig = go.FigureWidget(\n                layout=go.Layout(\n                    barmode=\"overlay\",\n                    margin={\"l\": 0, \"r\": 0, \"t\": 0, \"b\": 0},\n                    title=view_plane.value,\n                ),\n            )\n        else:\n            plotting_backends = [member.value for member in list(PlottingBackend)]\n            msg = f\"Plotting backend ({plotting_backend}) not supported, choose one of ({','.join(plotting_backends)}).\"\n            raise ValueError(msg)\n\n        self.slice_slider = widgets.IntSlider(\n            value=0,\n            min=0,\n            max=0,\n            description=\"Slice Index\",\n            continuous_update=False,\n        )\n        self.interactive = widgets.interactive_output(self.update, {\"slice_ix\": self.slice_slider})\n        self.children = [self.interactive, self.slice_slider]\n\n    @property\n    def cmap(self):\n        return self.parent.cmap\n\n    @property\n    def vol(self):\n        return self.parent.vol\n\n    @property\n    def spacing(self):\n        return self.parent.spacing\n\n    def update(self, slice_ix: int):\n        if self.vol is None:\n            self.fig.show()\n            return\n\n        slicing = self.view_plane.get_slice(slice_ix)\n        img = self.vol[slicing]\n\n        norm = Normalize(img.min(), img.max())\n        img = norm(img)\n        img = self.cmap(img)\n        img = (img * 255).astype(int)\n\n        x_axis, y_axis = self.view_plane.axes\n        dy, dx = self.spacing[x_axis], self.spacing[y_axis]\n\n        if self.plotting_backend == PlottingBackend.MPL:\n            aspect = dy / dx\n            if self.im is None or (self.im is not None and self.im.get_size()[:1] != img.shape[:1]):\n                self.im = self.ax.imshow(img, aspect=aspect)\n            else:\n                self.im.set_data(img)\n                self.ax.set_aspect(aspect)\n            display(self.fig)\n        elif self.plotting_backend == PlottingBackend.PLOTLY:\n            if not self.fig.data:\n                self.fig.add_image(z=img, dx=dx, dy=dy)\n                self.fig.show()\n                return\n\n            fig_img = self.fig.data[0]\n            fig_img.z = img\n            fig_img.dx = dx\n            fig_img.dy = dy\n\n            self.fig.show()\n        else:\n            plotting_backends = [member.value for member in list(PlottingBackend)]\n            msg = f\"Plotting backend ({plotting_backend}) not supported, choose one of ({','.join(plotting_backends)}).\"\n            raise ValueError(msg)\n\n\nclass ViewPlaneListWidget(widgets.HBox):\n    def __init__(self, parent, plotting_backend: PlottingBackend):\n        super().__init__()\n        self.parent = parent\n        self.children = [\n            ViewPlaneWidget(self, col, view_plane, plotting_backend) for col, view_plane in enumerate(ViewPlane, 1)\n        ]\n\n    @property\n    def vol(self):\n        return self.parent.vol\n\n    @property\n    def spacing(self):\n        return self.parent.spacing\n\n    @property\n    def cmap(self):\n        return self.parent.cmap\n\n\nclass DICOMViewerWidget(widgets.VBox):\n    def __init__(self, plotting_backend: PlottingBackend | str):\n        super().__init__()\n\n        if isinstance(plotting_backend, str):\n            plotting_backend = PlottingBackend(plotting_backend)\n\n        self.dataset_dropdown = widgets.Dropdown(\n            options=list(Dataset),\n            default=Dataset.TRAIN,\n            description=\"Dataset\",\n        )\n        self.study_id_dropdown = widgets.Dropdown(\n            options=[\"\", *fs.core.get_study_ids(train=True)],\n            description=\"Study ID\",\n            disabled=False,\n            default=\"\",\n        )\n        self.series_id_dropdown = widgets.Dropdown(\n            options=[\"\"],\n            description=\"Series ID\",\n            disabled=False,\n            default=\"\",\n        )\n\n        self.dataset_dropdown.observe(self.handle_dataset_change, names=\"value\")\n        self.study_id_dropdown.observe(self.handle_study_id_change, names=\"value\")\n        self.series_id_dropdown.observe(self.handle_series_id_change, names=\"value\")\n\n        self.vol = None\n        self.view_planes = ViewPlaneListWidget(parent=self, plotting_backend=plotting_backend)\n        self.children = [\n            self.dataset_dropdown,\n            self.study_id_dropdown,\n            self.series_id_dropdown,\n            self.view_planes,\n        ]\n        self.cmap = cm.get_cmap(\"bone\")\n\n    def handle_dataset_change(self, change: dict[str, Any]) -> None:\n        if change[\"new\"]:\n            train = Dataset(self.dataset_dropdown.value).is_train\n            study_ids = fs.core.get_study_ids(train=train)\n        else:\n            study_ids = []\n\n        self.study_id_dropdown.value = \"\"\n        self.study_id_dropdown.options = [\"\", *study_ids]\n\n    def handle_study_id_change(self, change: dict[str, Any]) -> None:\n        if change[\"new\"]:\n            train = Dataset(self.dataset_dropdown.value).is_train\n            series_ids = fs.core.get_series_ids(study_id=change[\"new\"], train=train)\n        else:\n            series_ids = []\n\n        self.series_id_dropdown.value = \"\"\n        self.series_id_dropdown.options = [\"\", *series_ids]\n\n    def handle_series_id_change(self, change: dict[str, Any]) -> None:\n        if not change[\"new\"]:\n            return\n\n        train = Dataset(self.dataset_dropdown.value).is_train\n        dicom_tags_df = train_dicom_tags_df if train else test_dicom_tags_df\n        vol, spacing = fs.core.load_series(\n            study_id=self.study_id_dropdown.value,\n            series_id=change[\"new\"],\n            dicom_tags_df=dicom_tags_df,\n            train=train,\n        )\n\n        self.vol = vol\n        self.spacing = spacing\n\n        for view_plane_widget in self.view_planes.children:\n            projected_axis = view_plane_widget.view_plane.projected_axis\n            view_plane_widget.slice_slider.max = vol.shape[projected_axis]\n            view_plane_widget.slice_slider.value = view_plane_widget.slice_slider.max // 2","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-06-30T20:58:07.519412Z","iopub.execute_input":"2024-06-30T20:58:07.519757Z","iopub.status.idle":"2024-06-30T20:58:07.646604Z","shell.execute_reply.started":"2024-06-30T20:58:07.519730Z","shell.execute_reply":"2024-06-30T20:58:07.645518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"w = DICOMViewerWidget(plotting_backend=Config.plotting_backend)\nw","metadata":{"tags":[],"execution":{"iopub.status.busy":"2024-06-30T20:59:47.388239Z","iopub.execute_input":"2024-06-30T20:59:47.388656Z","iopub.status.idle":"2024-06-30T20:59:47.531848Z","shell.execute_reply.started":"2024-06-30T20:59:47.388623Z","shell.execute_reply":"2024-06-30T20:59:47.530577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}