{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":100132,"sourceType":"modelInstanceVersion","modelInstanceId":64905,"modelId":89293},{"sourceId":111113,"sourceType":"modelInstanceVersion","modelInstanceId":85952,"modelId":84065}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Dependencies\n### timm_3d\nTimm models adapted for 3D data\n### torchio\nMedical imaging focused library. Used for 3D augmentations here\n### spacecutter\nOrdinal regression layer and related loss and callback functions i.e. `LogisticCumulativeLink`\n### open3d\nUsed for point cloud initialization and transformation. At first I used this for voxelization too, but dropped it in favor of a simpler sampling approach due to slow runtime.","metadata":{}},{"cell_type":"code","source":"%pip install --quiet /kaggle/input/timm_3d_deps/other/initial/9/pydicom/pydicom/pydicom-2.4.4-py3-none-any.whl\n%pip install timm_3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/timm_3d/\n%pip install torchio --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/torchio/\n%pip install itk --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/itk/itk\n%pip install skorch --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/skorch/skorch\n%pip install spacecutter --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/spacecutter/\n%pip install open3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/9/open3d","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:35:11.581854Z","iopub.execute_input":"2024-11-19T09:35:11.582102Z","iopub.status.idle":"2024-11-19T09:37:18.828729Z","shell.execute_reply.started":"2024-11-19T09:35:11.582075Z","shell.execute_reply":"2024-11-19T09:37:18.827809Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\"","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:18.830465Z","iopub.execute_input":"2024-11-19T09:37:18.830770Z","iopub.status.idle":"2024-11-19T09:37:18.835154Z","shell.execute_reply.started":"2024-11-19T09:37:18.830741Z","shell.execute_reply":"2024-11-19T09:37:18.834351Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndef retrieve_test_data(data_path):\n    test_df = pd.read_csv(data_path + 'test_series_descriptions.csv')\n\n    return test_df\n\nretrieve_test_data(data_path)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-19T09:37:18.836165Z","iopub.execute_input":"2024-11-19T09:37:18.836420Z","iopub.status.idle":"2024-11-19T09:37:19.181636Z","shell.execute_reply.started":"2024-11-19T09:37:18.836381Z","shell.execute_reply":"2024-11-19T09:37:19.180871Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndef retrieve_image_paths(base_path, study_id, series_id):\n    series_dir = os.path.join(base_path, str(study_id), str(series_id))\n    images = os.listdir(series_dir)\n    image_paths = [os.path.join(series_dir, img) for img in images]\n    return image_paths","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:19.184180Z","iopub.execute_input":"2024-11-19T09:37:19.184642Z","iopub.status.idle":"2024-11-19T09:37:19.189000Z","shell.execute_reply.started":"2024-11-19T09:37:19.184597Z","shell.execute_reply":"2024-11-19T09:37:19.188161Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loading\n\nQuite a few things are happening here. Refer here for more details on where the affine transforms come from: https://nipy.org/nibabel/dicom/dicom_orientation.html\n\nThe rough idea is -- converting the slices into point cloud format, transforming them into the same patient coordinate space, then converting into a voxel grid.\n\n1. Resizing each slice -- doing this first reduces the runtime from having to transform each point just to resize/downsample later anyway.\n2. Duplicating the slices and applying the affine transforms -- I found duplicating each slice by slice thickness makes it easier for the model to learn. The spatial information between slices might get somewhat warped, but the relative ordering is still the same but with much less empty space.\n3. Sampling into a voxel grid -- I am simply scaling the coordinates into the desired grid size, then sampling each channel by max.","metadata":{}},{"cell_type":"code","source":"import open3d as o3d\nfrom pydicom import dcmread\nimport math\nimport numpy as np\nimport cv2\nimport copy\n\ndef read_study_as_pcd(dir_path, series_types_dict=None, downsampling_factor=1, img_size=(256, 256)):\n    pcd_overall = o3d.geometry.PointCloud()\n\n    for path in glob.glob(os.path.join(dir_path, \"**/*.dcm\"), recursive=True):\n        dicom_slice = dcmread(path)\n\n        series_id = os.path.basename(os.path.dirname(path))\n        study_id = os.path.basename(os.path.dirname(os.path.dirname(path)))\n        if series_types_dict is None or int(series_id) not in series_types_dict:\n            series_desc = dicom_slice.SeriesDescription\n        else:\n            series_desc = series_types_dict[int(series_id)]\n            series_desc = series_desc.split(\" \")[-1]\n\n        x_orig, y_orig = dicom_slice.pixel_array.shape\n        img = np.expand_dims(cv2.resize(dicom_slice.pixel_array, img_size, interpolation=cv2.INTER_AREA), -1)\n        x, y, z = np.where(img)\n\n        downsampling_factor_iter = max(downsampling_factor, int(math.ceil(len(x) / 6e6)))\n\n        index_voxel = np.vstack((x, y, z))[:, ::downsampling_factor_iter]\n        grid_index_array = index_voxel.T\n        pcd = o3d.geometry.PointCloud(o3d.utility.Vector3dVector(grid_index_array.astype(np.float64)))\n\n        vals = np.expand_dims(img[x, y, z][::downsampling_factor_iter], -1)\n        if series_desc == \"T1\":\n            vals = np.pad(vals, ((0, 0), (0, 2)))\n        elif series_desc == \"T2\":\n            vals = np.pad(vals, ((0, 0), (1, 1)))\n        elif series_desc == \"T2/STIR\":\n            vals = np.pad(vals, ((0, 0), (2, 0)))\n        else:\n            raise ValueError(f\"Unknown series desc: {series_desc}\")\n\n        pcd.colors = o3d.utility.Vector3dVector(vals.astype(np.float64))\n\n        dX, dY = dicom_slice.PixelSpacing\n        dZ = dicom_slice.SliceThickness\n\n        X = np.array(list(dicom_slice.ImageOrientationPatient[:3]) + [0]) * dX\n        Y = np.array(list(dicom_slice.ImageOrientationPatient[3:]) + [0]) * dY\n\n        for z in range(int(dZ)):\n            pos = list(dicom_slice.ImagePositionPatient)\n            if series_desc == \"T2\":\n                pos[-1] += z\n            else:\n                pos[0] += z\n            S = np.array(pos + [1])\n\n            transform_matrix = np.array([X, Y, np.zeros(len(X)), S]).T\n            transform_matrix = transform_matrix @ np.matrix(\n                [[0, y_orig / img_size[1], 0, 0],\n                 [x_orig / img_size[0], 0, 0, 0],\n                 [0, 0, 1, 0],\n                 [0, 0, 0, 1]]\n            )\n\n            pcd_overall += copy.deepcopy(pcd).transform(transform_matrix)\n\n    return pcd_overall\n\n\n\ndef read_study_as_voxel_grid(dir_path, series_type_dict=None, downsampling_factor=1, img_size=(256, 256)):\n    pcd_overall = read_study_as_pcd(dir_path,\n                                    series_types_dict=series_type_dict,\n                                    downsampling_factor=downsampling_factor,\n                                    img_size=img_size)\n    box = pcd_overall.get_axis_aligned_bounding_box()\n\n    max_b = np.array(box.get_max_bound())\n    min_b = np.array(box.get_min_bound())\n\n    pts = (np.array(pcd_overall.points) - (min_b)) * (\n                (img_size[0] - 1, img_size[0] - 1, img_size[0] - 1) / (max_b - min_b))\n    coords = np.round(pts).astype(np.int32)\n    vals = np.array(pcd_overall.colors, dtype=np.float16)\n\n    grid = np.zeros((3, img_size[0], img_size[0], img_size[0]), dtype=np.float16)\n    indices = coords[:, 0], coords[:, 1], coords[:, 2]\n\n    np.maximum.at(grid[0], indices, vals[:, 0])\n    np.maximum.at(grid[1], indices, vals[:, 1])\n    np.maximum.at(grid[2], indices, vals[:, 2])\n\n\n    return grid","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:19.190246Z","iopub.execute_input":"2024-11-19T09:37:19.190701Z","iopub.status.idle":"2024-11-19T09:37:21.317099Z","shell.execute_reply.started":"2024-11-19T09:37:19.190663Z","shell.execute_reply":"2024-11-19T09:37:21.316373Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nimport torchio as tio\nimport torch.nn as nn\nimport pydicom\n\nCONDITIONS = {\n    \"Sagittal T2/STIR\": [\"Spinal Canal Stenosis\"],\n    \"Axial T2\": [\"Left Subarticular Stenosis\", \"Right Subarticular Stenosis\"],\n    \"Sagittal T1\": [\"Left Neural Foraminal Narrowing\", \"Right Neural Foraminal Narrowing\"],\n}\n\n\nclass PatientLevelTestset(Dataset):\n    def __init__(self,\n                 base_path: str,\n                 dataframe: pd.DataFrame,\n                 transform_3d=None):\n        self.base_path = base_path\n\n        self.dataframe = (dataframe[['study_id', \"series_id\", \"series_description\"]]\n                          .drop_duplicates())\n\n        self.subjects = self.dataframe[['study_id']].drop_duplicates().reset_index(drop=True)\n        self.series_descs = {e[0]: e[1] for e in self.dataframe[[\"series_id\", \"series_description\"]].drop_duplicates().values}\n\n        self.transform_3d = transform_3d\n\n    def __len__(self):\n        return len(self.subjects)\n\n    def __getitem__(self, index):\n        curr = self.subjects.iloc[index]\n        study_path = os.path.join(self.base_path, str(curr[\"study_id\"]))\n\n        study_images = read_study_as_voxel_grid(study_path, self.series_descs)\n\n        if self.transform_3d is not None:\n            study_images = self.transform_3d(torch.FloatTensor(study_images))  # .data\n            return study_images.to(torch.half), str(curr[\"study_id\"])\n\n        return torch.HalfTensor(study_images), str(curr[\"study_id\"])\n","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:21.318249Z","iopub.execute_input":"2024-11-19T09:37:21.318777Z","iopub.status.idle":"2024-11-19T09:37:25.000444Z","shell.execute_reply.started":"2024-11-19T09:37:21.318738Z","shell.execute_reply":"2024-11-19T09:37:24.999547Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform_3d = tio.Compose([\n    tio.RescaleIntensity([0, 1]),\n])","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:25.001688Z","iopub.execute_input":"2024-11-19T09:37:25.002236Z","iopub.status.idle":"2024-11-19T09:37:25.006804Z","shell.execute_reply.started":"2024-11-19T09:37:25.002194Z","shell.execute_reply":"2024-11-19T09:37:25.005942Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_subject_level_testset_and_loader(df: pd.DataFrame,\n                                             transform_3d,\n                                             base_path: str,\n                                             batch_size=1,\n                                             num_workers=0):\n    testset = PatientLevelTestset(base_path, df, transform_3d=transform_3d)\n    test_loader = DataLoader(testset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n\n    return testset, test_loader","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:25.008026Z","iopub.execute_input":"2024-11-19T09:37:25.008334Z","iopub.status.idle":"2024-11-19T09:37:29.008077Z","shell.execute_reply.started":"2024-11-19T09:37:25.008306Z","shell.execute_reply":"2024-11-19T09:37:29.007211Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndata = retrieve_test_data(data_path)\ndataset, dataloader = create_subject_level_testset_and_loader(data, transform_3d, os.path.join(data_path, \"test_images\"))","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:29.009197Z","iopub.execute_input":"2024-11-19T09:37:29.009571Z","iopub.status.idle":"2024-11-19T09:37:29.039272Z","shell.execute_reply.started":"2024-11-19T09:37:29.009533Z","shell.execute_reply":"2024-11-19T09:37:29.038678Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport glob \nimport torch\n\ngrid = dataset[0][0]\n\nfig, axs = plt.subplots(3, 3)\n\naxs[0, 0].imshow(grid[0, 128])\naxs[1, 0].imshow(grid[1, 128])\naxs[2, 0].imshow(grid[2, 128])\n\naxs[0, 1].imshow(grid[0, :, 128])\naxs[1, 1].imshow(grid[1, :, 128])\naxs[2, 1].imshow(grid[2, :, 128])\n\naxs[0, 2].imshow(grid[0, :, :, 128])\naxs[1, 2].imshow(grid[1, :, :, 128])\naxs[2, 2].imshow(grid[2, :, :, 128])\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-19T09:37:29.042727Z","iopub.execute_input":"2024-11-19T09:37:29.042959Z","iopub.status.idle":"2024-11-19T09:37:45.123691Z","shell.execute_reply.started":"2024-11-19T09:37:29.042936Z","shell.execute_reply":"2024-11-19T09:37:45.122870Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ndevice = torch.device(\"cuda\") if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:37:45.124755Z","iopub.execute_input":"2024-11-19T09:37:45.125010Z","iopub.status.idle":"2024-11-19T09:37:45.130483Z","shell.execute_reply.started":"2024-11-19T09:37:45.124983Z","shell.execute_reply":"2024-11-19T09:37:45.129665Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architecture\n\nI have somewhat better performance from a ViT. Might be due to better handling the gaps between slices. \n\nAlso note the LogisticCumulativeLink as the final head layer. This allows learning a continuous severity feature from ordinal labels and vice versa for inference.\n\n![](https://www.ethanrosenthal.com/2018/12/06/spacecutter-ordinal-regression/index_11_0.png)","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport timm_3d\nfrom spacecutter import *\nfrom spacecutter.losses import *\nfrom spacecutter.models import *\nfrom spacecutter.callbacks import *\n\n\nclass CNN_Model_3D_Multihead(nn.Module):\n    def __init__(self,\n                 backbone=\"efficientnet_lite0\",\n                 in_chans=1,\n                 out_classes=5,\n                 cutpoint_margin=0.15,\n                 pretrained=False):\n        super(CNN_Model_3D_Multihead, self).__init__()\n        self.out_classes = out_classes\n\n        self.encoder = timm_3d.create_model(\n            backbone,\n            features_only=False,\n            drop_rate=0,\n            drop_path_rate=0,\n            pretrained=pretrained,\n            in_chans=in_chans,\n            global_pool=\"max\"\n        )\n        if \"efficientnet\" in backbone:\n            head_in_dim = self.encoder.classifier.in_features\n            self.encoder.classifier = nn.Sequential(\n                nn.LayerNorm(head_in_dim),\n                nn.Dropout(0),\n            )\n\n        elif \"vit\" in backbone:\n            self.encoder.head.drop = nn.Dropout(0)\n            head_in_dim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n        self.heads = nn.ModuleList(\n            [nn.Sequential(\n                nn.Linear(head_in_dim, 1),\n                LogisticCumulativeLink(3)\n            ) for i in range(out_classes)]\n        )\n\n        self.ascension_callback = AscensionCallback(margin=cutpoint_margin)\n\n    def forward(self, x):\n        feat = self.encoder(x)\n        return torch.swapaxes(torch.stack([head(feat) for head in self.heads]), 0, 1)\n\n    def _ascension_callback(self):\n        for head in self.heads:\n            self.ascension_callback.clip(head[-1])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:37:45.131764Z","iopub.execute_input":"2024-11-19T09:37:45.132346Z","iopub.status.idle":"2024-11-19T09:37:47.444829Z","shell.execute_reply.started":"2024-11-19T09:37:45.132305Z","shell.execute_reply":"2024-11-19T09:37:47.444126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_Model_3D_Multihead(backbone=\"maxvit_rmlp_tiny_rw_256\", in_chans=3, out_classes=25).to(device)\nmodel.load_state_dict(torch.load(\"/kaggle/input/rsna-2024/pytorch/vit_voxel_v2/6/maxvit_rmlp_tiny_rw_256_256_v2_fold_3_32.pt\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:37:47.445923Z","iopub.execute_input":"2024-11-19T09:37:47.446248Z","iopub.status.idle":"2024-11-19T09:37:49.910937Z","shell.execute_reply.started":"2024-11-19T09:37:47.446213Z","shell.execute_reply":"2024-11-19T09:37:49.910123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CONDITIONS = {\n    \"Sagittal T2/STIR\": [\"spinal_canal_stenosis\"],\n    \"Axial T2\": [\"left_subarticular_stenosis\", \"right_subarticular_stenosis\"],\n    \"Sagittal T1\": [\"left_neural_foraminal_narrowing\", \"right_neural_foraminal_narrowing\"],\n}\n\nALL_CONDITIONS = sorted([\"spinal_canal_stenosis\", \"left_subarticular_stenosis\", \"right_subarticular_stenosis\", \"left_neural_foraminal_narrowing\", \"right_neural_foraminal_narrowing\"])\nLEVELS = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n\nresults_df = pd.DataFrame({\"row_id\":[], \"normal_mild\": [], \"moderate\": [], \"severe\": []})\n\nALL_CONDITIONS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:37:49.912007Z","iopub.execute_input":"2024-11-19T09:37:49.912360Z","iopub.status.idle":"2024-11-19T09:37:49.919761Z","shell.execute_reply.started":"2024-11-19T09:37:49.912319Z","shell.execute_reply":"2024-11-19T09:37:49.918942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Pre-populate results df\nimport glob\nimport os\n\nstudy_ids = glob.glob(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/*\")\nstudy_ids = [os.path.basename(e) for e in study_ids]\n\nresults_df = pd.DataFrame({\"row_id\":[], \"normal_mild\": [], \"moderate\": [], \"severe\": []})\nfor study_id in study_ids:\n    for condition in ALL_CONDITIONS:\n        for level in LEVELS:\n            row_id = f\"{study_id}_{condition}_{level}\"\n            results_df = results_df._append({\"row_id\": row_id, \"normal_mild\": 1/3, \"moderate\": 1/3, \"severe\": 1/3}, ignore_index=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:43:43.047754Z","iopub.execute_input":"2024-11-19T09:43:43.048120Z","iopub.status.idle":"2024-11-19T09:44:31.677427Z","shell.execute_reply.started":"2024-11-19T09:43:43.048090Z","shell.execute_reply":"2024-11-19T09:44:31.676446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset.dataframe","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:44:37.817381Z","iopub.execute_input":"2024-11-19T09:44:37.817776Z","iopub.status.idle":"2024-11-19T09:44:37.826334Z","shell.execute_reply.started":"2024-11-19T09:44:37.817743Z","shell.execute_reply":"2024-11-19T09:44:37.825371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.cuda.amp import autocast\nimport time\n\nstart_time = time.time()\n\nwith torch.no_grad():\n    with autocast(dtype=torch.float16):\n        model.eval()\n\n        for images, study_id in dataloader:\n            output = model(images.to(device))\n            for i, batch_out in enumerate(output):\n                batch_out = output.cpu().numpy()[i]\n                for index, level in enumerate(batch_out):\n                    row_id = f\"{study_id[i]}_{ALL_CONDITIONS[index // 5]}_{LEVELS[index % 5]}\"\n                    results_df.loc[results_df.row_id == row_id,'normal_mild'] = level[0]\n                    results_df.loc[results_df.row_id == row_id,'moderate'] = level[1]\n                    results_df.loc[results_df.row_id == row_id,'severe'] = level[2]\n                \nprint(\"--- %s seconds ---\" % (time.time() - start_time))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T09:44:39.933227Z","iopub.execute_input":"2024-11-19T09:44:39.934087Z","iopub.status.idle":"2024-11-19T09:44:54.372810Z","shell.execute_reply.started":"2024-11-19T09:44:39.934049Z","shell.execute_reply":"2024-11-19T09:44:54.371790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T10:09:52.422654Z","iopub.execute_input":"2024-11-19T10:09:52.423429Z","iopub.status.idle":"2024-11-19T10:09:52.436289Z","shell.execute_reply.started":"2024-11-19T10:09:52.423397Z","shell.execute_reply":"2024-11-19T10:09:52.435333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Load ground truth labels\nground_truth = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\n\n# Reshape the ground_truth to align with the results_df structure\nground_truth = ground_truth.melt(id_vars=[\"study_id\"], var_name=\"condition_level\", value_name=\"ground_truth_label\")\nground_truth['row_id'] = ground_truth['study_id'].astype(str) + \"_\" + ground_truth['condition_level']\n\n# Function to map predicted probabilities to labels\ndef map_prediction_to_label(row):\n    probabilities = [row['normal_mild'], row['moderate'], row['severe']]\n    max_prob_index = probabilities.index(max(probabilities))\n    return [\"Normal/Mild\", \"Moderate\", \"Severe\"][max_prob_index]\n\n# Apply the mapping function to results_df\nresults_df['predicted_label'] = results_df.apply(map_prediction_to_label, axis=1)\n\n# Merge results_df with ground truth on row_id\nmerged_df = results_df.merge(ground_truth[['row_id', 'ground_truth_label']], on='row_id', how='left')\n\n# Calculate accuracy\ncorrect_predictions = (merged_df['predicted_label'] == merged_df['ground_truth_label']).sum()\ntotal_predictions = len(merged_df)\naccuracy = correct_predictions / total_predictions\n\nprint(f\"Test Accuracy: {accuracy * 100:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T10:22:19.731370Z","iopub.execute_input":"2024-11-19T10:22:19.732180Z","iopub.status.idle":"2024-11-19T10:22:19.825073Z","shell.execute_reply.started":"2024-11-19T10:22:19.732144Z","shell.execute_reply":"2024-11-19T10:22:19.824193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ground_truth.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T10:27:24.850225Z","iopub.execute_input":"2024-11-19T10:27:24.850600Z","iopub.status.idle":"2024-11-19T10:27:24.859869Z","shell.execute_reply.started":"2024-11-19T10:27:24.850567Z","shell.execute_reply":"2024-11-19T10:27:24.858980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"merged_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T10:12:49.605226Z","iopub.execute_input":"2024-11-19T10:12:49.605581Z","iopub.status.idle":"2024-11-19T10:12:49.616298Z","shell.execute_reply.started":"2024-11-19T10:12:49.605548Z","shell.execute_reply":"2024-11-19T10:12:49.615426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix\nimport numpy as np\n\n# Ensure that both 'ground_truth_label' and 'predicted_label' are strings\nmerged_df['ground_truth_label'] = merged_df['ground_truth_label'].astype(str)\nmerged_df['predicted_label'] = merged_df['predicted_label'].astype(str)\n\n# Define the label categories\nlabels = [\"Normal/Mild\", \"Moderate\", \"Severe\"]\n\n# Generate the confusion matrix\nconf_matrix = confusion_matrix(merged_df['ground_truth_label'], merged_df['predicted_label'], labels=labels)\n\n# Create a heatmap using the confusion matrix\nplt.figure(figsize=(10, 7))\nsns.heatmap(conf_matrix, annot=True, fmt='d', cmap='Blues', xticklabels=labels, yticklabels=labels)\nplt.xlabel('Predicted Label')\nplt.ylabel('Actual Label')\nplt.title('Confusion Matrix - Prediction vs Ground Truth')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T10:10:03.425095Z","iopub.execute_input":"2024-11-19T10:10:03.425699Z","iopub.status.idle":"2024-11-19T10:10:04.172792Z","shell.execute_reply.started":"2024-11-19T10:10:03.425664Z","shell.execute_reply":"2024-11-19T10:10:04.171562Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}