{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":83273,"sourceType":"modelInstanceVersion","modelInstanceId":69945,"modelId":95069}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from glob import glob\nimport os\nfrom pathlib import Path\nimport re\n\nimport albumentations as A\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport pydicom\nimport torch\nfrom torch.utils.data import DataLoader, Dataset\nimport torch.nn as nn\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:39.071151Z","iopub.execute_input":"2024-07-27T22:07:39.071783Z","iopub.status.idle":"2024-07-27T22:07:42.111060Z","shell.execute_reply.started":"2024-07-27T22:07:39.071738Z","shell.execute_reply":"2024-07-27T22:07:42.110255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/input/2d_convnet/pytorch/default/1/2d_convnet.pth\"\ndata_dir = Path(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\")","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:42.112681Z","iopub.execute_input":"2024-07-27T22:07:42.113108Z","iopub.status.idle":"2024-07-27T22:07:42.117605Z","shell.execute_reply.started":"2024-07-27T22:07:42.113082Z","shell.execute_reply":"2024-07-27T22:07:42.116558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_weights(m):\n    if isinstance(m, nn.Conv2d):\n        nn.init.xavier_uniform_(m.weight)\n        if m.bias is not None:\n            nn.init.constant_(m.bias, 0)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:42.118806Z","iopub.execute_input":"2024-07-27T22:07:42.119187Z","iopub.status.idle":"2024-07-27T22:07:42.127282Z","shell.execute_reply.started":"2024-07-27T22:07:42.119154Z","shell.execute_reply":"2024-07-27T22:07:42.126335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:42.128329Z","iopub.execute_input":"2024-07-27T22:07:42.128650Z","iopub.status.idle":"2024-07-27T22:07:42.191507Z","shell.execute_reply.started":"2024-07-27T22:07:42.128626Z","shell.execute_reply":"2024-07-27T22:07:42.190599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvNet(nn.Module):\n    def __init__(self, num_classes=3, num_conditions=25):\n        super(ConvNet, self).__init__()\n        self.features = nn.Sequential(\n            # Conv Block 1\n            nn.Conv2d(30, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            \n            # Conv Block 2\n            nn.Conv2d(64, 128, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(128, 128, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            \n            # Conv Block 3\n            nn.Conv2d(128, 256, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, 256, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            \n            # Conv Block 4\n            nn.Conv2d(256, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=2, stride=2),\n            \n            # Conv Block 5\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(512, 512, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=2, stride=2)\n        )\n        \n        self.classifier = nn.Sequential(\n            nn.Linear(512 * 16 * 16, 4096),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p=0.5),\n            nn.Linear(4096, 4096),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p=0.5),\n            nn.Linear(4096, 1024),\n            nn.ReLU(inplace=True),\n            nn.Dropout(p=0.5)\n        )\n        \n        self.heads = nn.ModuleList([nn.Linear(1024, num_classes) for _ in range(num_conditions)])\n\n        # Apply Xavier Uniform initialization\n        self.apply(init_weights)\n  \n    def forward(self, x):\n        x = self.features(x)\n        x = x.view(x.size(0), -1)\n        x = self.classifier(x)\n        out = [head(x) for head in self.heads]\n        return torch.stack(out, dim=1)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:42.195233Z","iopub.execute_input":"2024-07-27T22:07:42.195616Z","iopub.status.idle":"2024-07-27T22:07:42.215249Z","shell.execute_reply.started":"2024-07-27T22:07:42.195575Z","shell.execute_reply":"2024-07-27T22:07:42.214282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ConvNet(num_classes=3, num_conditions=25)\nmodel.load_state_dict(torch.load(MODEL_PATH))\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:42.216655Z","iopub.execute_input":"2024-07-27T22:07:42.216952Z","iopub.status.idle":"2024-07-27T22:07:51.100253Z","shell.execute_reply.started":"2024-07-27T22:07:42.216924Z","shell.execute_reply":"2024-07-27T22:07:51.099405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    return [ atoi(c) for c in re.split(r'(\\d+)', text) ]","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.101552Z","iopub.execute_input":"2024-07-27T22:07:51.102216Z","iopub.status.idle":"2024-07-27T22:07:51.107929Z","shell.execute_reply.started":"2024-07-27T22:07:51.102181Z","shell.execute_reply":"2024-07-27T22:07:51.106939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dir = data_dir.joinpath(\"train_images\")\ntest_dir = data_dir.joinpath(\"test_images\")\ntest_series_descriptions_path = data_dir.joinpath(\"test_series_descriptions.csv\")\ntrain_series_descriptions_path = data_dir.joinpath(\"train_series_descriptions.csv\")\ntrain_path = data_dir.joinpath(\"train.csv\")\ntrain_labels_path = data_dir.joinpath(\"train_label_coordinates.csv\")\nsample_submission_path = data_dir.joinpath(\"sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.109250Z","iopub.execute_input":"2024-07-27T22:07:51.109632Z","iopub.status.idle":"2024-07-27T22:07:51.117690Z","shell.execute_reply.started":"2024-07-27T22:07:51.109600Z","shell.execute_reply":"2024-07-27T22:07:51.116906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(train_path)\ntest_descriptions_df = pd.read_csv(test_series_descriptions_path)\ntrain_descriptions_df = pd.read_csv(train_series_descriptions_path)\ntrain_labels_df = pd.read_csv(train_labels_path)\nsample_submission_df = pd.read_csv(sample_submission_path)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.118693Z","iopub.execute_input":"2024-07-27T22:07:51.118945Z","iopub.status.idle":"2024-07-27T22:07:51.222352Z","shell.execute_reply.started":"2024-07-27T22:07:51.118920Z","shell.execute_reply":"2024-07-27T22:07:51.221599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_CHANNELS = 30\nIMG_SIZE = [512, 512]\nAUG_PROB = 0.75","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.223438Z","iopub.execute_input":"2024-07-27T22:07:51.223764Z","iopub.status.idle":"2024-07-27T22:07:51.228317Z","shell.execute_reply.started":"2024-07-27T22:07:51.223741Z","shell.execute_reply":"2024-07-27T22:07:51.227323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_val = A.Compose([\n#     A.Resize(IMG_SIZE, IMG_SIZE),\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.229628Z","iopub.execute_input":"2024-07-27T22:07:51.229978Z","iopub.status.idle":"2024-07-27T22:07:51.237738Z","shell.execute_reply.started":"2024-07-27T22:07:51.229946Z","shell.execute_reply":"2024-07-27T22:07:51.236738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNA24TestDataset(Dataset):\n    def __init__(self, df, study_ids, phase='test', transform=transforms_val):\n        self.df = df\n        self.study_ids = study_ids\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.df[self.df['study_id']==study_id]\n        pdf_ = pdf[pdf['series_description']==series_desc]\n        allimgs = []\n        for i, row in pdf_.iterrows():\n            # pimgs = data_dir.joinpath(\"test_images\", str(study_id), str(row[\"series_id\"])).glob(\"*.dcm\")\n            pimgs = glob(f'{data_dir}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm')\n            pimgs = sorted(pimgs, key=natural_keys)\n            allimgs.extend(pimgs)\n            \n        return allimgs\n    \n    def read_dcm_ret_arr(self, src_path):\n        dicom_data = pydicom.dcmread(src_path)\n        image = dicom_data.pixel_array\n        image = (image - image.min()) / (image.max() - image.min() + 1e-6) * 255\n        img = cv2.resize(image, (IMG_SIZE[0], IMG_SIZE[1]),interpolation=cv2.INTER_CUBIC)\n        assert img.shape==(IMG_SIZE[0], IMG_SIZE[1])\n        return img\n\n    def __getitem__(self, idx):\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], INPUT_CHANNELS), dtype=np.uint8)\n        st_id = self.study_ids[idx]        \n        \n        # Sagittal T1\n        allimgs_st1 = self.get_img_paths(st_id, 'Sagittal T1')\n        if len(allimgs_st1)==0:\n            print(st_id, ': Sagittal T1, has no images')\n        \n        else:\n            step = len(allimgs_st1) / 10.0\n            st = len(allimgs_st1) / 2.0 - 4.0 * step\n            end = len(allimgs_st1) + 0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i - 0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st1[ind2])\n                    x[..., j] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T1')\n                    pass\n            \n        # Sagittal T2/STIR\n        allimgs_st2 = self.get_img_paths(st_id, 'Sagittal T2/STIR')\n        if len(allimgs_st2)==0:\n            print(st_id, ': Sagittal T2/STIR, has no images')\n            \n        else:\n            step = len(allimgs_st2) / 10.0\n            st = len(allimgs_st2) / 2.0 - 4.0 * step\n            end = len(allimgs_st2) + 0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i - 0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st2[ind2])\n                    x[..., j+10] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                    pass\n            \n        # Axial T2\n        allimgs_at2 = self.get_img_paths(st_id, 'Axial T2')\n        if len(allimgs_at2)==0:\n            print(st_id, ': Axial T2, has no images')\n            \n        else:\n            step = len(allimgs_at2) / 10.0\n            st = len(allimgs_at2) / 2.0 - 4.0 * step\n            end = len(allimgs_at2) + 0.0001\n\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i - 0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_at2[ind2])\n                    x[..., j+20] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Axial T2')\n                    pass              \n            \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)\n                \n        return torch.tensor(x, dtype=torch.float32), str(st_id)\n#         return torch.tensor(x, dtype=torch.float32).unsqueeze(0),","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.239237Z","iopub.execute_input":"2024-07-27T22:07:51.239583Z","iopub.status.idle":"2024-07-27T22:07:51.261955Z","shell.execute_reply.started":"2024-07-27T22:07:51.239544Z","shell.execute_reply":"2024-07-27T22:07:51.260885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids = list(test_descriptions_df['study_id'].unique())\nstudy_ids","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.263116Z","iopub.execute_input":"2024-07-27T22:07:51.263389Z","iopub.status.idle":"2024-07-27T22:07:51.276221Z","shell.execute_reply.started":"2024-07-27T22:07:51.263366Z","shell.execute_reply":"2024-07-27T22:07:51.275185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RSNA24TestDataset(test_descriptions_df, study_ids, transform=transforms_val)\ntest_dl = DataLoader(\n    test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=0,\n    pin_memory=True,\n    drop_last=False\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.280425Z","iopub.execute_input":"2024-07-27T22:07:51.281191Z","iopub.status.idle":"2024-07-27T22:07:51.286158Z","shell.execute_reply.started":"2024-07-27T22:07:51.281156Z","shell.execute_reply":"2024-07-27T22:07:51.285311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = [\n    'spinal_canal_stenosis', \n    'left_neural_foraminal_narrowing', \n    'right_neural_foraminal_narrowing',\n    'left_subarticular_stenosis',\n    'right_subarticular_stenosis'\n]\n\nLEVELS = [\n    'l1_l2',\n    'l2_l3',\n    'l3_l4',\n    'l4_l5',\n    'l5_s1',\n]","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.287093Z","iopub.execute_input":"2024-07-27T22:07:51.287375Z","iopub.status.idle":"2024-07-27T22:07:51.297118Z","shell.execute_reply.started":"2024-07-27T22:07:51.287352Z","shell.execute_reply":"2024-07-27T22:07:51.296162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# autocast = torch.cuda.amp.autocast(enabled=True, dtype=torch.half)\nNUM_CONDITIONS = 25\ny_preds = []\nrow_names = []\nmodel.eval()\nwith tqdm(test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (x, si) in enumerate(pbar):\n            x = x.to(device)\n            pred_per_study = np.zeros((25, 3))\n            \n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(si[0] + '_' + cond + '_' + level)\n            \n#             with autocast:\n            y = model(x)[0]\n            y_preds = y.float().softmax(1).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:51.298295Z","iopub.execute_input":"2024-07-27T22:07:51.298782Z","iopub.status.idle":"2024-07-27T22:07:52.427868Z","shell.execute_reply.started":"2024-07-27T22:07:51.298749Z","shell.execute_reply":"2024-07-27T22:07:52.426929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_preds","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:52.429141Z","iopub.execute_input":"2024-07-27T22:07:52.429479Z","iopub.status.idle":"2024-07-27T22:07:52.437280Z","shell.execute_reply.started":"2024-07-27T22:07:52.429451Z","shell.execute_reply":"2024-07-27T22:07:52.436200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABELS = list(sample_submission_df.columns[1:])\n# LABELS\nsample_submission_df","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:52.438808Z","iopub.execute_input":"2024-07-27T22:07:52.439206Z","iopub.status.idle":"2024-07-27T22:07:52.460083Z","shell.execute_reply.started":"2024-07-27T22:07:52.439171Z","shell.execute_reply":"2024-07-27T22:07:52.458952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub.head(100)","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:52.461450Z","iopub.execute_input":"2024-07-27T22:07:52.461933Z","iopub.status.idle":"2024-07-27T22:07:52.483467Z","shell.execute_reply.started":"2024-07-27T22:07:52.461886Z","shell.execute_reply":"2024-07-27T22:07:52.482230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"execution":{"iopub.status.busy":"2024-07-27T22:07:52.485092Z","iopub.execute_input":"2024-07-27T22:07:52.485573Z","iopub.status.idle":"2024-07-27T22:07:52.502954Z","shell.execute_reply.started":"2024-07-27T22:07:52.485514Z","shell.execute_reply":"2024-07-27T22:07:52.501800Z"},"trusted":true},"execution_count":null,"outputs":[]}]}