{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9114148,"sourceType":"datasetVersion","datasetId":5501147}],"dockerImageVersionId":30761,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\nimport sys\nsys.path.append('/kaggle/input/timm-3d/')\n\nimport os\nimport glob\nimport pydicom\nimport cv2\nimport random\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport timm_3d","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-09-19T09:33:19.923514Z","iopub.execute_input":"2024-09-19T09:33:19.925493Z","iopub.status.idle":"2024-09-19T09:33:19.937698Z","shell.execute_reply.started":"2024-09-19T09:33:19.925432Z","shell.execute_reply":"2024-09-19T09:33:19.935711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# simplified, ref from: https://www.kaggle.com/code/hengck23/ver-2-more-magic-single-stage-model","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:33:19.940753Z","iopub.execute_input":"2024-09-19T09:33:19.941220Z","iopub.status.idle":"2024-09-19T09:33:19.956830Z","shell.execute_reply.started":"2024-09-19T09:33:19.941181Z","shell.execute_reply":"2024-09-19T09:33:19.955456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_to_8bit(x):\n    lower, upper = np.percentile(x, (1, 99))\n    x = np.clip(x, lower, upper)\n    x = x - np.min(x)\n    x = x / np.max(x)\n    return (x * 255).astype(\"uint8\")\n\n\ndef load_dicom_stack(dicom_folder, plane, reverse_sort=False, dicom_files=None,\n                     img_size=384):\n    if dicom_files is None:\n        dicom_files = glob.glob(os.path.join(dicom_folder, \"*.dcm\"))\n    dicoms = [pydicom.dcmread(f) for f in dicom_files]\n    dicom_instance_numbers = [int(i.split('/')[-1][:-4]) for i in dicom_files]\n\n    plane = {\"sagittal\": 0, \"coronal\": 1, \"axial\": 2}[plane.lower()]\n    positions = np.asarray([float(d.ImagePositionPatient[plane]) for d in dicoms])\n    # if reverse_sort=False, then increasing array index will be from RIGHT->LEFT and CAUDAL->CRANIAL\n    # thus we do reverse_sort=True for axial so increasing array index is craniocaudal\n    idx = np.argsort(-positions if reverse_sort else positions)\n    ipp = np.asarray([d.ImagePositionPatient for d in dicoms]).astype(\"float\")[idx]\n    # array = np.stack([d.pixel_array.astype(\"float32\") for d in dicoms])\n\n    dicom_instance_numbers = np.array(dicom_instance_numbers)[idx]\n    cols = np.asarray([d.pixel_array.shape[1] for d in dicoms]).astype(\"int\")[idx].tolist()\n    rows = np.asarray([d.pixel_array.shape[0] for d in dicoms]).astype(\"int\")[idx].tolist()\n    instance_num_to_shape = {}\n    dicom_instance_numbers = dicom_instance_numbers.tolist()\n    for i, n in enumerate(dicom_instance_numbers):\n        instance_num_to_shape[n] = [cols[i], rows[i]]\n\n    dicom_instance_numbers_to_idx = {}\n    for i, ins in enumerate(dicom_instance_numbers):\n        dicom_instance_numbers_to_idx[ins] = i\n\n    array = []\n    for i, d in enumerate(dicoms):\n        a = d.pixel_array.astype(\"float32\")\n        a = torch.from_numpy(a).unsqueeze(0).unsqueeze(0)\n        a = F.interpolate(a, (img_size, img_size)).numpy().squeeze()\n        array.append(a)\n    array = np.array(array)\n    array = array[idx]\n\n    return {\"array\": convert_to_8bit(array), \"positions\": ipp,\n            \"pixel_spacing\": np.asarray(dicoms[0].PixelSpacing).astype(\"float\"),\n            \"instance_num_to_shape\": instance_num_to_shape,\n            \"dicom_instance_numbers_to_idx\": dicom_instance_numbers_to_idx\n            }","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:33:19.958989Z","iopub.execute_input":"2024-09-19T09:33:19.959509Z","iopub.status.idle":"2024-09-19T09:33:19.977369Z","shell.execute_reply.started":"2024-09-19T09:33:19.959436Z","shell.execute_reply":"2024-09-19T09:33:19.975769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self,\n                 data_dir,\n                 df,\n                 phase='train',\n                 series_description='Sagittal T2/STIR',\n                 data_cache_dir='./cache2/',\n                 ):\n        super(RSNADataset, self).__init__()\n        self.phase = phase\n        self.data_cache_dir = data_cache_dir\n        self.data_dir = data_dir\n        desc_df = pd.read_csv(f\"{data_dir}/train_series_descriptions.csv\")\n        self.image_dir = f\"{data_dir}/train_images/\"\n        self.desc_df = desc_df[desc_df['series_description'].isin([series_description])]\n        self.coord_df = pd.read_csv(f\"{data_dir}/train_label_coordinates.csv\")\n        level2id = {\n            'L1/L2': 0,\n            'L2/L3': 1,\n            'L3/L4': 2,\n            'L4/L5': 3,\n            'L5/S1': 4\n        }\n        self.coord_df = self.coord_df.replace(level2id)\n        self.phase = phase\n        self.img_size = 192\n        self.depth = 32\n\n        self.series_description = series_description\n        self.keypoints_dict = {}\n\n        study_ids = []\n        for _, row in df.iterrows():\n            study_id = row['study_id']\n            g = self.desc_df[self.desc_df['study_id'] == study_id]\n            series_id_list = g['series_id'].to_list()\n            if len(series_id_list) == 0:\n                print(f'{study_id} has no {series_description} series')\n            infos = []\n            for sid in series_id_list:\n                coord_sub_df = self.coord_df[self.coord_df['study_id'] == study_id]\n                coord_sub_df = coord_sub_df[coord_sub_df['series_id'] == sid]\n                if series_description == 'Sagittal T2/STIR':\n                    keypoints = -np.ones((5, 1, 3), dtype=np.float32)\n                else:\n                    keypoints = -np.ones((5, 2, 3), dtype=np.float32)\n\n                for _, row in coord_sub_df.iterrows():\n                    idx = 0\n                    if row['condition'] in \\\n                            ['Right Neural Foraminal Narrowing',\n                             'Right Subarticular Stenosis']:\n                        idx = 1\n                    x, y, instance_number = row['x'], row['y'], row['instance_number']\n                    keypoints[row['level'], idx, 0] = x\n                    keypoints[row['level'], idx, 1] = y\n                    keypoints[row['level'], idx, 2] = instance_number\n\n                keypoints = keypoints.transpose(1, 0, 2)\n                infos.append({\n                    'series_id': sid,\n                    'keypoints': keypoints.reshape(-1, 3),\n                })\n            if len(infos) > 0:\n                self.keypoints_dict[study_id] = infos\n                study_ids.append(study_id)\n        self.df = df[df['study_id'].isin(study_ids)]\n\n    def __len__(self):\n        return len(self.df)\n\n    # keypoints_norm (2, 5, 3)\n    # grade_labels (2, 5)\n    def make_grade_mask(self, keypoints_norm, grade_labels,\n                        mask_h=80, mask_w=80, mask_depth=32):\n        keypoints_norm = keypoints_norm.reshape(-1, 3)\n        grade_labels = grade_labels.reshape(-1)\n        N = grade_labels.shape[0]\n        mask = np.zeros((N, mask_depth, mask_h, mask_w), dtype=np.int8)\n        radius = 3  # you can tune this\n\n        for i in range(N):\n            p = keypoints_norm[i]\n            la = grade_labels[i]\n            x, y, z = p\n            x, y, z = int(np.round(x)), int(np.round(y)), int(np.round(z))\n            if z > 0:\n                z0 = z - 1  # you can tune this\n                if z0 < 0:\n                    z0 = 0\n                z1 = z + 1  # you can tune this\n                if z1 > mask_depth - 1:\n                    z1 = mask_depth - 1\n                for iz in range(z0, z1):\n                    la = int(la)\n                    mask[i, iz] = cv2.circle(mask[i, iz], (x, y), radius, la + 1, -1, cv2.LINE_4)\n            else:\n                mask[i, :] = -100  # ignore index\n        return mask\n\n    def make_volume_and_mask(self, study_id, series_id, keypoints, label):\n        fn1 = f'{self.data_cache_dir}/{study_id}_{series_id}_volume.npz'\n        fn2 = f'{self.data_cache_dir}/{study_id}_{series_id}_keypoints.npy'\n        if os.path.exists(fn1):\n            arr = np.load(fn1)['arr_0']\n            keypoints = np.load(fn2)\n        else:\n\n            dicom_dir = f'{self.data_dir}/train_images/{study_id}/{series_id}/'\n\n            reverse_sort = False\n            if self.series_description == 'Axial T2':\n                reverse_sort = True\n            dicom = load_dicom_stack(dicom_dir,\n                                     plane='sagittal',\n                                     reverse_sort=reverse_sort,\n                                     img_size=self.img_size)\n            instance_num_to_shape = dicom[\"instance_num_to_shape\"]\n            dicom_instance_numbers_to_idx = dicom[\"dicom_instance_numbers_to_idx\"]\n\n            arr = dicom['array']\n            for i in range(len(keypoints)):\n                x, y, ins_num = keypoints[i]\n                # no labeled\n                if x < 0:\n                    continue\n                origin_w, origin_h = instance_num_to_shape[ins_num]\n                x = self.img_size / origin_w * x\n                y = self.img_size / origin_h * y\n                z = dicom_instance_numbers_to_idx[ins_num]\n                keypoints[i, 0] = x\n                keypoints[i, 1] = y\n                keypoints[i, 2] = z\n            if os.path.exists(self.data_cache_dir):\n                np.savez_compressed(fn1, arr)\n                np.save(fn2, keypoints)\n\n        keypoints_xy = keypoints[:, :2]\n        keypoints_z = keypoints[:, 2:]\n\n        mask_xy = np.where(keypoints_xy < 0, 0, 1)\n        mask_z = np.where(keypoints_z < 0, 0, 1)\n        keypoints_xy = keypoints_xy\n        # print('keypoints_xy shape: ', keypoints_xy.shape)\n\n        keypoints_xy[:, :2] = keypoints_xy[:, :2]\n\n        arr = arr.astype(np.float32)\n        arr = torch.from_numpy(arr).unsqueeze(0).unsqueeze(0)\n        arr = F.interpolate(arr, size=(self.depth, self.img_size, self.img_size)).squeeze(0)\n\n        keypoints = np.concatenate((keypoints_xy, keypoints_z), axis=1)\n\n        label_mask = self.make_grade_mask(keypoints,\n                                          label,\n                                          mask_h=self.img_size,\n                                          mask_w=self.img_size,\n                                          mask_depth=self.depth)\n        return arr, keypoints, label_mask\n\n    def __getitem__(self, idx):\n        item = self.df.iloc[idx]\n        study_id = int(item['study_id'])\n        label = item[1:].values.astype(np.int64)\n        if self.series_description == 'Sagittal T2/STIR':\n            label = label.reshape(5, 5)[0:1].reshape(-1)\n        elif self.series_description == 'Sagittal T1':\n            label = label.reshape(5, 5)[1:3].reshape(-1)\n        else:\n            label = label.reshape(5, 5)[3:].reshape(-1)\n\n        info_list = self.keypoints_dict[study_id]\n        if True:\n            info = info_list[0]  # random.choice(info_list)\n            keypoints = info['keypoints']\n            series_id = info['series_id']\n            arr, keypoints, label_mask = self.make_volume_and_mask(\n                study_id, series_id, keypoints, label\n            )\n            keypoints = torch.from_numpy(keypoints).float()\n            label_mask = torch.from_numpy(label_mask).long()\n\n            label = torch.from_numpy(label).long()\n            return {\n                'imgs': arr,\n                'keypoints': keypoints.reshape(-1),\n                'label_mask': label_mask,\n                'label': label,\n                'study_id': study_id,\n                'series_id': series_id,\n            }","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:33:20.054977Z","iopub.execute_input":"2024-09-19T09:33:20.055521Z","iopub.status.idle":"2024-09-19T09:33:20.095043Z","shell.execute_reply.started":"2024-09-19T09:33:20.055473Z","shell.execute_reply":"2024-09-19T09:33:20.093477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dense_label_mask_to_sparse(dense_label_mask, label):\n    n, d, h, w = dense_label_mask.shape\n    device = dense_label_mask.device\n    sparse_mask = torch.zeros(n, 3, d, h, w, device=device)\n    sparse_mask[:] = -1e8\n    for i in range(n):\n        la = int(label[i])\n        # skip < 0\n        if la >= 0:\n            dense = dense_label_mask[i]\n            sparse_mask[i, la] = torch.where(dense > 0,\n                                             torch.full(dense.shape, 1e8, device=device),\n                                             torch.full(dense.shape, -1e8, device=device))\n      \n    return sparse_mask\n\n\ndef sparse_mask_to_grade(sparse_mask):\n    grade = torch.sum(sparse_mask, dim=(2, 3, 4))\n    return grade\n\n\ndef sparse_mask_to_xyz(sparse_mask):\n    num_point, num_grade3, D, H, W = sparse_mask.shape\n    prob = sparse_mask.flatten(1).softmax(-1).reshape(num_point, num_grade3, D, H, W)\n    device = prob.device\n    x = torch.linspace(0, W - 1, W, device=device)\n    y = torch.linspace(0, H - 1, H, device=device)\n    z = torch.linspace(0, D - 1, D, device=device)\n    pos_x = x.reshape(1, 1, 1, 1, W)\n    pos_y = y.reshape(1, 1, 1, H, 1)\n    pos_z = z[:D].reshape(1, 1, D, 1, 1)\n\n    px = torch.sum(pos_x * prob, dim=(1, 2, 3, 4)).unsqueeze(1)\n    py = torch.sum(pos_y * prob, dim=(1, 2, 3, 4)).unsqueeze(1)\n    pz = torch.sum(pos_z * prob, dim=(1, 2, 3, 4)).unsqueeze(1)\n    xyz = torch.cat((px, py, pz), dim=1)\n\n    return xyz","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:33:20.098169Z","iopub.execute_input":"2024-09-19T09:33:20.099350Z","iopub.status.idle":"2024-09-19T09:33:20.114582Z","shell.execute_reply.started":"2024-09-19T09:33:20.099290Z","shell.execute_reply":"2024-09-19T09:33:20.112887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if True:\n    import tqdm\n    import albumentations as A\n\n    data_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    df = pd.read_csv(f'{data_dir}/train.csv')\n    df = df[:100]\n    df = df.fillna(-100)\n    label2id = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n    df = df.replace(label2id)\n    \n    dset = RSNADataset(data_dir,\n                     df,\n                     phase='train',\n                     series_description='Sagittal T2/STIR',\n                     data_cache_dir='./cache2/')\n    dataloader = DataLoader(dset, num_workers=4, batch_size=4)\n    weights = torch.tensor([1.0, 2.0, 4.0])\n    ce_loss_fn = nn.CrossEntropyLoss(weight=weights)\n    keypoint_loss_fn = nn.L1Loss()\n\n    pixel_err_list = []\n    ce_loss_list = []\n    for t in tqdm.tqdm(dataloader):\n        imgs = t['imgs']\n\n        bs = imgs.shape[0]\n        keypoints = t['keypoints'].reshape(bs, -1, 3)\n\n        label = t['label'].reshape(bs, -1)\n        label_mask = t['label_mask']\n\n        all_xyz = []\n        all_grade = []\n        all_sparse_mask = []\n        for b in range(bs):\n            sparse_mask = dense_label_mask_to_sparse(label_mask[b], label[b])\n            xyz = sparse_mask_to_xyz(sparse_mask)\n            grade = sparse_mask_to_grade(sparse_mask)\n            all_xyz.append(xyz.unsqueeze(0))\n            all_grade.append(grade.unsqueeze(0))\n            all_sparse_mask.append(sparse_mask.unsqueeze(0))\n        all_xyz = torch.cat(all_xyz)\n        all_grade = torch.cat(all_grade).reshape(-1, 3)\n\n        label = label.reshape(-1)\n        pixel_err = keypoint_loss_fn(keypoints, all_xyz)\n        ce_loss = ce_loss_fn(all_grade, label)\n\n        pixel_err_list.append(pixel_err.item())\n        ce_loss_list.append(ce_loss.item())\n        \n\n    pixel_err = np.mean(pixel_err_list)\n    ce_loss = np.mean(ce_loss_list)\n    print('pixel_err: ', pixel_err)\n    print('ce loss: ', ce_loss)","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:33:20.116590Z","iopub.execute_input":"2024-09-19T09:33:20.117121Z","iopub.status.idle":"2024-09-19T09:34:00.262690Z","shell.execute_reply.started":"2024-09-19T09:33:20.117067Z","shell.execute_reply":"2024-09-19T09:34:00.260816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# with  model here: https://www.kaggle.com/code/hengck23/ver-2-more-magic-single-stage-model","metadata":{"execution":{"iopub.status.busy":"2024-09-19T09:34:00.265629Z","iopub.execute_input":"2024-09-19T09:34:00.266167Z","iopub.status.idle":"2024-09-19T09:34:00.273828Z","shell.execute_reply.started":"2024-09-19T09:34:00.266122Z","shell.execute_reply":"2024-09-19T09:34:00.271510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FPN_3D_Head4(nn.Module):\n    def __init__(self,\n                 channels=[64, 256, 512, 1024],\n                 out_c=256):\n        super(FPN_3D_Head4, self).__init__()\n        self.conv1 = nn.Conv3d(channels[0], out_c, kernel_size=(1, 1, 1))\n        self.conv2 = nn.Conv3d(channels[1], out_c, kernel_size=(1, 1, 1))\n        self.conv3 = nn.Conv3d(channels[2], out_c, kernel_size=(1, 1, 1))\n        self.conv4 = nn.Conv3d(channels[3], out_c, kernel_size=(1, 1, 1))\n\n    def _upsample_add(self, x, y):\n        _, _, H, W, D = y.size()\n        return F.interpolate(x, size=(H, W, D), mode='trilinear') + y\n\n    def forward(self, x1, x2, x3, x4):\n        p4 = self.conv4(x4)\n        p3 = self._upsample_add(p4, self.conv3(x3))\n        p2 = self._upsample_add(p3, self.conv2(x2))\n        p1 = self._upsample_add(p2, self.conv1(x1))\n        _, _, H, W, D = p1.size()\n        p1 = F.interpolate(p1, size=(H, W, 4 * D), mode='trilinear')\n        return p1\n\n\nclass Model3D(nn.Module):\n    def __init__(self,\n                 model_name='densenet161',\n                 n_grade=5,\n                 pretrained=False):\n        super().__init__()\n        self.model_name = model_name\n        self.backbone = timm_3d.create_model(\n            model_name,\n            pretrained=pretrained,\n            features_only=True,\n            in_chans=1,\n            global_pool='none',\n        )\n        _backbone_fea_channels = {\n            'densenet161': [96, 384, 768, 2112, 2208],\n        }\n\n        if 'densenet' in model_name:\n            channels = np.array(_backbone_fea_channels[model_name])[1:].tolist()\n        else:\n            channels = _backbone_fea_channels[model_name]\n\n        self.fpn = FPN_3D_Head4(channels, out_c=64)\n        # self.backbone.features_conv0.stride = (2, 2, 1)\n        self.linear = nn.Conv3d(64, n_grade * 3, (1, 1, 1), (1, 1, 1))\n        self.n_grade = n_grade\n\n    def forward(self, x):\n        # b, c, d, h, w = x.shape\n        x = x.permute(0, 1, 3, 4, 2)\n        if 'densenet' in self.model_name:\n            _, x1, x2, x3, x4 = self.backbone(x)\n        else:\n            x1, x2, x3, x4 = self.backbone(x)\n        p = self.fpn(x1, x2, x3, x4)\n        p = p.permute(0, 1, 4, 2, 3)\n        p = self.linear(p)\n        b, c, d, h, w = p.shape\n        p = p.reshape(b, self.n_grade, 3, d, h, w)\n        return p\n","metadata":{},"execution_count":null,"outputs":[]}]}