{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9457725,"sourceType":"datasetVersion","datasetId":5749506}],"dockerImageVersionId":30762,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport warnings\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom PIL import Image\nimport pydicom\nimport cv2\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\nimport random\nfrom sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:02.033775Z","iopub.execute_input":"2024-09-22T19:53:02.034128Z","iopub.status.idle":"2024-09-22T19:53:30.973284Z","shell.execute_reply.started":"2024-09-22T19:53:02.034079Z","shell.execute_reply":"2024-09-22T19:53:30.972471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = 512\nin_channels = 30\nnum_classes = 75\nbatch_size = 16\nn_workers = os.cpu_count()\nmodel_name = 'tf_efficientnet_b3.ns_jft_in1k'\nlr = 1e-4\nepochs = 10\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\ncomm_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:30.974848Z","iopub.execute_input":"2024-09-22T19:53:30.975378Z","iopub.status.idle":"2024-09-22T19:53:31.028537Z","shell.execute_reply.started":"2024-09-22T19:53:30.975341Z","shell.execute_reply":"2024-09-22T19:53:31.027500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tdesc_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.030479Z","iopub.execute_input":"2024-09-22T19:53:31.030861Z","iopub.status.idle":"2024-09-22T19:53:31.059686Z","shell.execute_reply.started":"2024-09-22T19:53:31.030826Z","shell.execute_reply":"2024-09-22T19:53:31.058730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tdesc_df","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.061972Z","iopub.execute_input":"2024-09-22T19:53:31.062698Z","iopub.status.idle":"2024-09-22T19:53:31.082121Z","shell.execute_reply.started":"2024-09-22T19:53:31.062664Z","shell.execute_reply":"2024-09-22T19:53:31.081292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids = list(tdesc_df['study_id'].unique())","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.083232Z","iopub.execute_input":"2024-09-22T19:53:31.083509Z","iopub.status.idle":"2024-09-22T19:53:31.090446Z","shell.execute_reply.started":"2024-09-22T19:53:31.083477Z","shell.execute_reply":"2024-09-22T19:53:31.089514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.091553Z","iopub.execute_input":"2024-09-22T19:53:31.091885Z","iopub.status.idle":"2024-09-22T19:53:31.102622Z","shell.execute_reply.started":"2024-09-22T19:53:31.091839Z","shell.execute_reply":"2024-09-22T19:53:31.101680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_val = A.Compose([\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.103715Z","iopub.execute_input":"2024-09-22T19:53:31.104028Z","iopub.status.idle":"2024-09-22T19:53:31.112985Z","shell.execute_reply.started":"2024-09-22T19:53:31.103997Z","shell.execute_reply":"2024-09-22T19:53:31.112130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNATestDataset(Dataset):\n    def __init__(self, tdesc_df, study_ids, transform=None):\n        self.tdesc_df = tdesc_df\n        self.study_ids = study_ids\n        self.transforms = transform\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def __getitem__(self, idx):\n        X_test = np.zeros((img_size, img_size, in_channels), dtype=np.uint8)\n        st_id = self.study_ids[idx]\n        temp_df = tdesc_df[tdesc_df['study_id']==st_id]\n        for i, row in temp_df.iterrows():\n            study_id = row['study_id']\n            series_id = row['series_id']\n            series_desc = row['series_description']\n            if series_desc == 'Axial T2':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X_test[:, :, j] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Axial T2')\n                    pass \n                \n            if series_desc == 'Sagittal T1': \n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)                    \n                        X_test[:, :, j+10] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except Exception as e:\n                    print(f'failed to load on {study_id}, Sagittal T1, {e}')\n                    pass \n            \n            if series_desc == 'Sagittal T2/STIR':\n                try:\n                    j = 0\n                    for img_file in os.listdir(f'{comm_path}/test_images/{study_id}/{series_id}'):\n                        src_path  = f'{comm_path}/test_images/{study_id}/{series_id}/{img_file}'\n                        dcm_data = pydicom.dcmread(src_path)\n                        img = dcm_data.pixel_array\n                        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255\n                        img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_CUBIC)\n                        img = np.array(img)           \n                        X_test[:, :, j+20] = img.astype(np.uint8)\n                        j=j+1\n                        if j==10:\n                            break\n                except:\n                    print(f'failed to load on {study_id}, Sagittal T2/STIR')\n                    pass \n        \n        \n        X_test = self.transforms(image=X_test)[\"image\"]\n        X_test = X_test.transpose(2, 0, 1)\n        \n        return X_test, st_id","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.114281Z","iopub.execute_input":"2024-09-22T19:53:31.114554Z","iopub.status.idle":"2024-09-22T19:53:31.132678Z","shell.execute_reply.started":"2024-09-22T19:53:31.114523Z","shell.execute_reply":"2024-09-22T19:53:31.131856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RSNATestDataset(tdesc_df, study_ids, transform=transforms_val)\ntest_dl = DataLoader(test_ds, batch_size=1, shuffle=False, pin_memory=True, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.133667Z","iopub.execute_input":"2024-09-22T19:53:31.133949Z","iopub.status.idle":"2024-09-22T19:53:31.144945Z","shell.execute_reply.started":"2024-09-22T19:53:31.133919Z","shell.execute_reply":"2024-09-22T19:53:31.144225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAModel(nn.Module):\n    def __init__(self, model_name, in_channels, num_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, in_chans=in_channels, num_classes=num_classes, global_pool='avg')\n    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.148244Z","iopub.execute_input":"2024-09-22T19:53:31.148538Z","iopub.status.idle":"2024-09-22T19:53:31.156672Z","shell.execute_reply.started":"2024-09-22T19:53:31.148505Z","shell.execute_reply":"2024-09-22T19:53:31.155830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nCKPT_PATHS = glob.glob('/kaggle/input/rsna-model-folds/model_fold*.pt')\nCKPT_PATHS = sorted(CKPT_PATHS)","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.157844Z","iopub.execute_input":"2024-09-22T19:53:31.158446Z","iopub.status.idle":"2024-09-22T19:53:31.173372Z","shell.execute_reply.started":"2024-09-22T19:53:31.158405Z","shell.execute_reply":"2024-09-22T19:53:31.172612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor i, cp in enumerate(CKPT_PATHS):\n    print(f'loading {cp}...')\n    model = RSNAModel(model_name, in_channels, num_classes, pretrained=False)\n    model.load_state_dict(torch.load(cp, weights_only=True))\n    model.eval()\n    model.half()\n    model.to(device)\n    models.append(model)","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:31.174562Z","iopub.execute_input":"2024-09-22T19:53:31.175425Z","iopub.status.idle":"2024-09-22T19:53:35.855813Z","shell.execute_reply.started":"2024-09-22T19:53:31.175392Z","shell.execute_reply":"2024-09-22T19:53:35.854735Z"},"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]\nrow_names = []\ny_preds = []\nautocast = torch.cuda.amp.autocast(enabled=True, dtype=torch.half)\nwith tqdm(test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (X, st_id) in enumerate(pbar):\n            sm = nn.Softmax(dim=1)\n            X = X.to(device)\n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(f'{st_id.item()}_{cond}_{level}')\n            with autocast:\n                y_pred = 0\n                for m in models:\n                    y = m(X)[0].reshape((25, 3))\n                    y = sm(y)\n                    y = y.cpu().numpy()\n                    y_pred += y\n                y_preds.append(y_pred / len(models))","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:35.857038Z","iopub.execute_input":"2024-09-22T19:53:35.857390Z","iopub.status.idle":"2024-09-22T19:53:38.187768Z","shell.execute_reply.started":"2024-09-22T19:53:35.857356Z","shell.execute_reply":"2024-09-22T19:53:38.186849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_preds = np.concatenate(y_preds, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.189040Z","iopub.execute_input":"2024-09-22T19:53:38.189403Z","iopub.status.idle":"2024-09-22T19:53:38.193907Z","shell.execute_reply.started":"2024-09-22T19:53:38.189369Z","shell.execute_reply":"2024-09-22T19:53:38.192916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv\")\nLABELS = list(sub_df.columns[1:])\nLABELS","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.195251Z","iopub.execute_input":"2024-09-22T19:53:38.195649Z","iopub.status.idle":"2024-09-22T19:53:38.214711Z","shell.execute_reply.started":"2024-09-22T19:53:38.195605Z","shell.execute_reply":"2024-09-22T19:53:38.213294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"row_names","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.216146Z","iopub.execute_input":"2024-09-22T19:53:38.216476Z","iopub.status.idle":"2024-09-22T19:53:38.223204Z","shell.execute_reply.started":"2024-09-22T19:53:38.216443Z","shell.execute_reply":"2024-09-22T19:53:38.222145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_preds","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.224406Z","iopub.execute_input":"2024-09-22T19:53:38.224707Z","iopub.status.idle":"2024-09-22T19:53:38.235378Z","shell.execute_reply.started":"2024-09-22T19:53:38.224674Z","shell.execute_reply":"2024-09-22T19:53:38.234441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame()\nsubmission['row_id'] = row_names\nsubmission[LABELS] = y_preds\nsubmission.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.236820Z","iopub.execute_input":"2024-09-22T19:53:38.237265Z","iopub.status.idle":"2024-09-22T19:53:38.259356Z","shell.execute_reply.started":"2024-09-22T19:53:38.237221Z","shell.execute_reply":"2024-09-22T19:53:38.258354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"execution":{"iopub.status.busy":"2024-09-22T19:53:38.260757Z","iopub.execute_input":"2024-09-22T19:53:38.261624Z","iopub.status.idle":"2024-09-22T19:53:38.284917Z","shell.execute_reply.started":"2024-09-22T19:53:38.261579Z","shell.execute_reply":"2024-09-22T19:53:38.284055Z"},"trusted":true},"execution_count":null,"outputs":[]}]}