{"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":9231709,"sourceType":"datasetVersion","datasetId":5583775}],"dockerImageVersionId":30762,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Import the Libraries","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import KFold\nfrom collections import OrderedDict\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 timm\nfrom timm.utils import ModelEmaV2\nfrom transformers import get_cosine_schedule_with_warmup\nimport albumentations as A\nfrom sklearn.model_selection import KFold\nimport re\nimport pydicom","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:12.464146Z","iopub.execute_input":"2024-09-09T18:13:12.464471Z","iopub.status.idle":"2024-09-09T18:13:41.519084Z","shell.execute_reply.started":"2024-09-09T18:13:12.464435Z","shell.execute_reply":"2024-09-09T18:13:41.518101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:41.520850Z","iopub.execute_input":"2024-09-09T18:13:41.521385Z","iopub.status.idle":"2024-09-09T18:13:42.044270Z","shell.execute_reply.started":"2024-09-09T18:13:41.521348Z","shell.execute_reply":"2024-09-09T18:13:42.043468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load & Edit The Data","metadata":{}},{"cell_type":"code","source":"rd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\nOUTPUT_DIR = f'/kaggle/input/rsna2024-lsdc-training-baseline/rsna24-results'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.045436Z","iopub.execute_input":"2024-09-09T18:13:42.045911Z","iopub.status.idle":"2024-09-09T18:13:42.091812Z","shell.execute_reply.started":"2024-09-09T18:13:42.045874Z","shell.execute_reply":"2024-09-09T18:13:42.090208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Configuration\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\nSEED = 8620\nIMG_SIZE = (512, 512)\nIN_CHANS = 42\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\nN_FOLDS = 5\nMODEL_NAME = \"edgenext_base.in21k_ft_in1k\"\nBATCH_SIZE = 1\ndata_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.096121Z","iopub.execute_input":"2024-09-09T18:13:42.096495Z","iopub.status.idle":"2024-09-09T18:13:42.105295Z","shell.execute_reply.started":"2024-09-09T18:13:42.096459Z","shell.execute_reply":"2024-09-09T18:13:42.104487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Device setup\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\n# Load DataFrames\ndf = pd.read_csv(f'{data_dir}/test_series_descriptions.csv')\nstudy_ids = df['study_id'].unique().tolist()\n\nsample_sub = pd.read_csv(f'{data_dir}/sample_submission.csv')\nLABELS = sample_sub.columns[1:].tolist()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.106394Z","iopub.execute_input":"2024-09-09T18:13:42.106740Z","iopub.status.idle":"2024-09-09T18:13:42.138786Z","shell.execute_reply.started":"2024-09-09T18:13:42.106706Z","shell.execute_reply":"2024-09-09T18:13:42.137978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Conditions and Levels\nCONDITIONS = [\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-09-09T18:13:42.139949Z","iopub.execute_input":"2024-09-09T18:13:42.140338Z","iopub.status.idle":"2024-09-09T18:13:42.145431Z","shell.execute_reply.started":"2024-09-09T18:13:42.140289Z","shell.execute_reply":"2024-09-09T18:13:42.144472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Helper functions\ndef 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-09-09T18:13:42.146741Z","iopub.execute_input":"2024-09-09T18:13:42.147163Z","iopub.status.idle":"2024-09-09T18:13:42.156632Z","shell.execute_reply.started":"2024-09-09T18:13:42.147120Z","shell.execute_reply":"2024-09-09T18:13:42.155785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Feature Engineering","metadata":{}},{"cell_type":"code","source":"import glob\nimport cv2\nimport pydicom\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\n\nclass 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        series_df = self.df[(self.df['study_id'] == study_id) & \n                            (self.df['series_description'] == series_desc)]\n        img_paths = []\n        for _, row in series_df.iterrows():\n            paths = sorted(glob.glob(f'{data_dir}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm'), \n                           key=natural_keys)\n            img_paths.extend(paths)\n        return img_paths\n\n    def read_dcm_image(self, src_path):\n        dicom_data = pydicom.dcmread(src_path)\n        image = dicom_data.pixel_array\n        norm_img = (image - image.min()) / (image.max() - image.min() + 1e-6) * 255\n        resized_img = cv2.resize(norm_img, IMG_SIZE, interpolation=cv2.INTER_CUBIC)\n        return resized_img.astype(np.uint8)\n\n    def load_series_images(self, study_id, series_desc, start_idx):\n        images = np.zeros((IMG_SIZE[0], IMG_SIZE[1], 14), dtype=np.uint8)\n        img_paths = self.get_img_paths(study_id, series_desc)\n        \n        if not img_paths:\n            print(f'{study_id}: {series_desc} has no images')\n            return images\n        \n        step = len(img_paths) / 14.0\n        mid_point = len(img_paths) / 2.0 - 6.0 * step\n        \n        for j, i in enumerate(np.arange(mid_point, len(img_paths), step)):\n            try:\n                idx = max(0, int(round(i - 0.5)))\n                images[..., j] = self.read_dcm_image(img_paths[idx])\n            except Exception as e:\n                print(f'Failed to load {series_desc} for {study_id}: {e}')\n                \n        return images\n\n    def __getitem__(self, idx):\n        study_id = self.study_ids[idx]\n        channels = 42\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], channels), dtype=np.uint8)\n\n        # Load images for each series\n        x[..., :14] = self.load_series_images(study_id, 'Sagittal T1', 0)\n        x[..., 14:28] = self.load_series_images(study_id, 'Sagittal T2/STIR', 14)\n        x[..., 28:] = self.load_series_images(study_id, 'Axial T2', 28)\n\n        # Apply transformations\n        if self.transform:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)  # Channels-first for PyTorch\n        return x, str(study_id)\n\n# Transformations\ntransforms_test = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])\n\n# Dataset and DataLoader\ntest_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\ntest_dl = DataLoader(test_ds, batch_size=1, shuffle=False, num_workers=N_WORKERS,\n                     pin_memory=True, drop_last=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.157726Z","iopub.execute_input":"2024-09-09T18:13:42.157982Z","iopub.status.idle":"2024-09-09T18:13:42.178831Z","shell.execute_reply.started":"2024-09-09T18:13:42.157953Z","shell.execute_reply":"2024-09-09T18:13:42.177767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport timm\n\nclass RSNA24Model(nn.Module):\n    \"\"\"\n    Custom model class for RSNA 2024 Lumbar Spine Degenerative Classification.\n    \n    Args:\n        model_name (str): Name of the model architecture from the TIMM library.\n        in_c (int): Number of input channels. Default is 42.\n        n_classes (int): Number of output classes. Default is 75.\n        pretrained (bool): If True, use pre-trained weights. Default is True.\n        features_only (bool): If True, return features only instead of final output. Default is False.\n    \"\"\"\n    \n    def __init__(self, model_name, in_c=42, n_classes=75, pretrained=True, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            features_only=features_only,\n            in_chans=in_c,\n            num_classes=n_classes,\n            global_pool='avg'\n        )\n    \n    def forward(self, x):\n        \"\"\"\n        Forward pass for the model.\n        \n        Args:\n            x (torch.Tensor): Input tensor with shape (batch_size, in_c, height, width).\n            \n        Returns:\n            torch.Tensor: Model output with shape (batch_size, n_classes).\n        \"\"\"\n        return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.180088Z","iopub.execute_input":"2024-09-09T18:13:42.180381Z","iopub.status.idle":"2024-09-09T18:13:42.195033Z","shell.execute_reply.started":"2024-09-09T18:13:42.180349Z","shell.execute_reply":"2024-09-09T18:13:42.194067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.198103Z","iopub.execute_input":"2024-09-09T18:13:42.198515Z","iopub.status.idle":"2024-09-09T18:13:42.218531Z","shell.execute_reply.started":"2024-09-09T18:13:42.198481Z","shell.execute_reply":"2024-09-09T18:13:42.217471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\n# Constants\nCKPT_PATHS = sorted([\n    \"/kaggle/input/rsna-2024-edgenext-base/model_fold-0.pt\",\n    \"/kaggle/input/rsna-2024-edgenext-base/model_fold-1.pt\",\n    \"/kaggle/input/rsna-2024-edgenext-base/model_fold-2.pt\",\n])\n\n# Load Models\nmodels = []\nfor cp in CKPT_PATHS:\n    print(f'Loading checkpoint: {cp}...')\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=False)\n    try:\n        model.load_state_dict(torch.load(cp))\n    except Exception as e:\n        print(f'Error loading {cp}: {e}')\n        continue\n\n    model.eval()\n    model.half()\n    model.to(device)\n    models.append(model)\n\n# Autocast for mixed precision\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\n\n# Predictions and Submission\ny_preds = []\nrow_names = []\n\nwith torch.no_grad():\n    for x, study_id in tqdm(test_dl, leave=True):\n        x = x.to(device)\n        pred_per_study = np.zeros((N_LABELS, 3))\n\n        # Generate row names for submission\n        for cond in CONDITIONS:\n            for level in LEVELS:\n                row_names.append(f'{study_id[0]}_{cond}_{level}')\n\n        with autocast:\n            for model in models:\n                y = model(x)[0]\n                for col in range(N_LABELS):\n                    pred = y[col * 3 : (col + 1) * 3]\n                    y_pred = pred.float().softmax(dim=0).cpu().numpy()\n                    pred_per_study[col] += y_pred / len(models)\n        \n        y_preds.append(pred_per_study)\n\n# Combine and save predictions\ny_preds = np.concatenate(y_preds, axis=0)\n\nsubmission_df = pd.DataFrame({\n    'row_id': row_names,\n    **{label: y_preds[:, i] for i, label in enumerate(LABELS)}\n})\n\nsubmission_df.to_csv('submission.csv', index=False)\n\n# Verify Submission\nprint(pd.read_csv('submission.csv').head(3))\n","metadata":{"execution":{"iopub.status.busy":"2024-09-09T18:13:42.219954Z","iopub.execute_input":"2024-09-09T18:13:42.220349Z","iopub.status.idle":"2024-09-09T18:13:48.361113Z","shell.execute_reply.started":"2024-09-09T18:13:42.220312Z","shell.execute_reply":"2024-09-09T18:13:48.360105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}