{"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":190051144,"sourceType":"kernelVersion"},{"sourceId":190053369,"sourceType":"kernelVersion"}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nfrom PIL import Image\nimport cv2\nimport math, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\n\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import AdamW\nimport torchvision.models as models\n\nimport timm\nfrom timm.utils import ModelEmaV2\nfrom transformers import ViTForImageClassification, ViTFeatureExtractor, ViTConfig, TrainingArguments, Trainer, DefaultDataCollator, get_cosine_schedule_with_warmup\n\nimport albumentations as A\n\nfrom sklearn.model_selection import KFold\n\nimport re\nimport pydicom","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-07-29T17:15:06.898681Z","iopub.execute_input":"2024-07-29T17:15:06.899056Z","iopub.status.idle":"2024-07-29T17:15:25.547995Z","shell.execute_reply.started":"2024-07-29T17:15:06.899003Z","shell.execute_reply":"2024-07-29T17:15:25.547219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:25.549512Z","iopub.execute_input":"2024-07-29T17:15:25.550110Z","iopub.status.idle":"2024-07-29T17:15:25.554465Z","shell.execute_reply.started":"2024-07-29T17:15:25.550082Z","shell.execute_reply":"2024-07-29T17:15:25.553487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/glsdc-resnet-train/rsna24-results","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:25.555519Z","iopub.execute_input":"2024-07-29T17:15:25.555764Z","iopub.status.idle":"2024-07-29T17:15:26.581894Z","shell.execute_reply.started":"2024-07-29T17:15:25.555743Z","shell.execute_reply":"2024-07-29T17:15:26.580688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"OUTPUT_DIR = f'/kaggle/input/rsna2024-lsdc-training-baseline/rsna24-results'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\nSEED = 8620\n\nIMG_SIZE = [512, 512]\nIN_CHANS = 30\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\n\nN_FOLDS = 5\n\nMODEL_NAME = \"tf_efficientnet_b3.ns_jft_in1k\"\n\nBATCH_SIZE = 1","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.585067Z","iopub.execute_input":"2024-07-29T17:15:26.585747Z","iopub.status.idle":"2024-07-29T17:15:26.638518Z","shell.execute_reply.started":"2024-07-29T17:15:26.585718Z","shell.execute_reply":"2024-07-29T17:15:26.637434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.639917Z","iopub.execute_input":"2024-07-29T17:15:26.640362Z","iopub.status.idle":"2024-07-29T17:15:26.649745Z","shell.execute_reply.started":"2024-07-29T17:15:26.640328Z","shell.execute_reply":"2024-07-29T17:15:26.648870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f'{rd}/test_series_descriptions.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.650828Z","iopub.execute_input":"2024-07-29T17:15:26.651411Z","iopub.status.idle":"2024-07-29T17:15:26.682990Z","shell.execute_reply.started":"2024-07-29T17:15:26.651388Z","shell.execute_reply":"2024-07-29T17:15:26.682202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids = list(df['study_id'].unique())","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.684050Z","iopub.execute_input":"2024-07-29T17:15:26.684365Z","iopub.status.idle":"2024-07-29T17:15:26.690711Z","shell.execute_reply.started":"2024-07-29T17:15:26.684343Z","shell.execute_reply":"2024-07-29T17:15:26.689660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub = pd.read_csv(f'{rd}/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.691883Z","iopub.execute_input":"2024-07-29T17:15:26.692165Z","iopub.status.idle":"2024-07-29T17:15:26.701935Z","shell.execute_reply.started":"2024-07-29T17:15:26.692140Z","shell.execute_reply":"2024-07-29T17:15:26.700935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LABELS = list(sample_sub.columns[1:])\nLABELS","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.703063Z","iopub.execute_input":"2024-07-29T17:15:26.703415Z","iopub.status.idle":"2024-07-29T17:15:26.709804Z","shell.execute_reply.started":"2024-07-29T17:15:26.703379Z","shell.execute_reply":"2024-07-29T17:15:26.708828Z"},"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-29T17:15:26.712975Z","iopub.execute_input":"2024-07-29T17:15:26.713484Z","iopub.status.idle":"2024-07-29T17:15:26.717986Z","shell.execute_reply.started":"2024-07-29T17:15:26.713458Z","shell.execute_reply":"2024-07-29T17:15:26.717170Z"},"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-29T17:15:26.719121Z","iopub.execute_input":"2024-07-29T17:15:26.719355Z","iopub.status.idle":"2024-07-29T17:15:26.726973Z","shell.execute_reply.started":"2024-07-29T17:15:26.719335Z","shell.execute_reply":"2024-07-29T17:15:26.726033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Dataset","metadata":{}},{"cell_type":"code","source":"class RSNA24TestDataset(Dataset):\n    def __init__(self, df, study_ids, phase='test', transform=None):\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 = glob.glob(f'{rd}/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], IN_CHANS), 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            \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 x, str(st_id)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.728457Z","iopub.execute_input":"2024-07-29T17:15:26.728831Z","iopub.status.idle":"2024-07-29T17:15:26.750125Z","shell.execute_reply.started":"2024-07-29T17:15:26.728802Z","shell.execute_reply":"2024-07-29T17:15:26.749274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_test = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.751244Z","iopub.execute_input":"2024-07-29T17:15:26.751575Z","iopub.status.idle":"2024-07-29T17:15:26.760419Z","shell.execute_reply.started":"2024-07-29T17:15:26.751545Z","shell.execute_reply":"2024-07-29T17:15:26.759613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\ntest_dl = DataLoader(\n    test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.761483Z","iopub.execute_input":"2024-07-29T17:15:26.761781Z","iopub.status.idle":"2024-07-29T17:15:26.768166Z","shell.execute_reply.started":"2024-07-29T17:15:26.761757Z","shell.execute_reply":"2024-07-29T17:15:26.767155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Model","metadata":{}},{"cell_type":"code","source":"models_list = []\n\nimport glob\nCKPT_PATHS = glob.glob('/kaggle/input/glsdc-resnet-train/rsna24-results/best_wll_model_fold-*.pt')\nCKPT_PATHS = sorted(CKPT_PATHS)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.769320Z","iopub.execute_input":"2024-07-29T17:15:26.769651Z","iopub.status.idle":"2024-07-29T17:15:26.777273Z","shell.execute_reply.started":"2024-07-29T17:15:26.769622Z","shell.execute_reply":"2024-07-29T17:15:26.776342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/glsdc-resnet-download/resnet50_modified.pth","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:26.778698Z","iopub.execute_input":"2024-07-29T17:15:26.779070Z","iopub.status.idle":"2024-07-29T17:15:27.784960Z","shell.execute_reply.started":"2024-07-29T17:15:26.779014Z","shell.execute_reply":"2024-07-29T17:15:27.784010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = '/kaggle/input/glsdc-resnet-download/resnet50_modified.pth'\n\nfor i, cp in enumerate(CKPT_PATHS):\n    print(f'Loading {cp}...')\n\n    # Create a new ResNet model instance\n    model = models.resnet50(pretrained=False)\n\n    # Modify the first convolutional layer to accept IN_CHANS input channels\n    model.conv1 = torch.nn.Conv2d(\n        in_channels=IN_CHANS,\n        out_channels=model.conv1.out_channels,\n        kernel_size=model.conv1.kernel_size,\n        stride=model.conv1.stride,\n        padding=model.conv1.padding,\n        bias=model.conv1.bias is not None\n    )\n\n    # Modify the final fully connected layer to output N_CLASSES\n    num_features = model.fc.in_features\n    model.fc = torch.nn.Linear(num_features, N_CLASSES)\n\n    # Load the model state dict from the checkpoint\n    model.load_state_dict(torch.load(model_path))\n\n    # Prepare the model for evaluation\n    model.eval()\n    model.half()  # Convert to half precision if needed\n    model.to(device)  # Move model to the appropriate device\n\n    models_list.append(model)\n\nprint(\"All models loaded and prepared.\")","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:27.786379Z","iopub.execute_input":"2024-07-29T17:15:27.786698Z","iopub.status.idle":"2024-07-29T17:15:31.914377Z","shell.execute_reply.started":"2024-07-29T17:15:27.786671Z","shell.execute_reply":"2024-07-29T17:15:31.913475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\ny_preds = []\nrow_names = []\n\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)  # Ensure data is on GPU\n            pred_per_study = np.zeros((N_LABELS, 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                x = x.half()  # Cast the input tensor to half precision (only on GPU)\n                for m in models_list:\n                    y = m(x)\n                    logits = y  # Assuming your model outputs the logits directly\n\n                    for col in range(N_LABELS):\n                        start_idx = col * 3\n                        end_idx = start_idx + 3\n                        pred = logits[:, start_idx:end_idx]  # Assuming logits shape is (batch_size, N_LABELS * 3)\n\n                        y_pred = pred.float().softmax(1).cpu().numpy()\n\n                        # Accumulate predictions across models\n                        pred_per_study[col] += y_pred[0]  # Assuming single batch size\n\n            # Average the predictions\n            pred_per_study /= len(models_list)\n            y_preds.append(pred_per_study)\n\ny_preds = np.concatenate(y_preds, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:31.915692Z","iopub.execute_input":"2024-07-29T17:15:31.915968Z","iopub.status.idle":"2024-07-29T17:15:34.146594Z","shell.execute_reply.started":"2024-07-29T17:15:31.915943Z","shell.execute_reply":"2024-07-29T17:15:34.145575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub.head(25)","metadata":{"execution":{"iopub.status.busy":"2024-07-29T17:15:57.075284Z","iopub.execute_input":"2024-07-29T17:15:57.076227Z","iopub.status.idle":"2024-07-29T17:15:57.103209Z","shell.execute_reply.started":"2024-07-29T17:15:57.076187Z","shell.execute_reply":"2024-07-29T17:15:57.101665Z"},"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-29T17:16:00.509780Z","iopub.execute_input":"2024-07-29T17:16:00.510374Z","iopub.status.idle":"2024-07-29T17:16:00.530231Z","shell.execute_reply.started":"2024-07-29T17:16:00.510341Z","shell.execute_reply":"2024-07-29T17:16:00.529431Z"},"trusted":true},"execution_count":null,"outputs":[]}]}