{"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"}],"dockerImageVersionId":30761,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"raw","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"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-23T22:57:25.157532Z","iopub.execute_input":"2024-09-23T22:57:25.158935Z","iopub.status.idle":"2024-09-23T22:57:25.167202Z","shell.execute_reply.started":"2024-09-23T22:57:25.158858Z","shell.execute_reply":"2024-09-23T22:57:25.165854Z"},"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-23T22:57:25.170126Z","iopub.execute_input":"2024-09-23T22:57:25.170659Z","iopub.status.idle":"2024-09-23T22:57:25.184728Z","shell.execute_reply.started":"2024-09-23T22:57:25.170602Z","shell.execute_reply":"2024-09-23T22:57:25.183401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-23T22:57:25.186143Z","iopub.execute_input":"2024-09-23T22:57:25.186650Z","iopub.status.idle":"2024-09-23T22:57:25.197577Z","shell.execute_reply.started":"2024-09-23T22:57:25.186592Z","shell.execute_reply":"2024-09-23T22:57:25.196247Z"},"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-23T22:57:25.200882Z","iopub.execute_input":"2024-09-23T22:57:25.201470Z","iopub.status.idle":"2024-09-23T22:57:25.210332Z","shell.execute_reply.started":"2024-09-23T22:57:25.201410Z","shell.execute_reply":"2024-09-23T22:57:25.209065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Device setup\n\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2024-09-23T22:57:25.212675Z","iopub.execute_input":"2024-09-23T22:57:25.213095Z","iopub.status.idle":"2024-09-23T22:57:25.222010Z","shell.execute_reply.started":"2024-09-23T22:57:25.213052Z","shell.execute_reply":"2024-09-23T22:57:25.220856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load DataFrames\ndf = pd.read_csv(f'{data_dir}/test_series_descriptions.csv')\nstudy_ids = df['study_id'].unique().tolist()\nsample_sub = pd.read_csv(f'{data_dir}/sample_submission.csv')\nLABELS = sample_sub.columns[1:].tolist()","metadata":{"execution":{"iopub.status.busy":"2024-09-23T22:57:25.223815Z","iopub.execute_input":"2024-09-23T22:57:25.224654Z","iopub.status.idle":"2024-09-23T22:57:25.242874Z","shell.execute_reply.started":"2024-09-23T22:57:25.224591Z","shell.execute_reply":"2024-09-23T22:57:25.241459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#  Conditions and Levels\n\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-23T22:57:25.244486Z","iopub.execute_input":"2024-09-23T22:57:25.244900Z","iopub.status.idle":"2024-09-23T22:57:25.251201Z","shell.execute_reply.started":"2024-09-23T22:57:25.244850Z","shell.execute_reply":"2024-09-23T22:57:25.250027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Helper functions\ndef atoi(text):\n    \n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    \n    return [atoi(c) for c in re.split(r'(\\d+)', text)]","metadata":{"execution":{"iopub.status.busy":"2024-09-23T22:57:25.252811Z","iopub.execute_input":"2024-09-23T22:57:25.253298Z","iopub.status.idle":"2024-09-23T22:57:25.263068Z","shell.execute_reply.started":"2024-09-23T22:57:25.253255Z","shell.execute_reply":"2024-09-23T22:57:25.261838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-23T22:57:25.317261Z","iopub.execute_input":"2024-09-23T22:57:25.318253Z","iopub.status.idle":"2024-09-23T22:57:25.342381Z","shell.execute_reply.started":"2024-09-23T22:57:25.318201Z","shell.execute_reply":"2024-09-23T22:57:25.341039Z"},"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)\n","metadata":{"execution":{"iopub.status.busy":"2024-09-23T22:57:25.344716Z","iopub.execute_input":"2024-09-23T22:57:25.345219Z","iopub.status.idle":"2024-09-23T22:57:25.358911Z","shell.execute_reply.started":"2024-09-23T22:57:25.345149Z","shell.execute_reply":"2024-09-23T22:57:25.357637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head()","metadata":{"execution":{"iopub.status.busy":"2024-09-23T22:57:25.360658Z","iopub.execute_input":"2024-09-23T22:57:25.361332Z","iopub.status.idle":"2024-09-23T22:57:25.380857Z","shell.execute_reply.started":"2024-09-23T22:57:25.361264Z","shell.execute_reply":"2024-09-23T22:57:25.379671Z"},"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-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-23T22:57:25.382429Z","iopub.execute_input":"2024-09-23T22:57:25.382903Z","iopub.status.idle":"2024-09-23T22:57:28.535284Z","shell.execute_reply.started":"2024-09-23T22:57:25.382856Z","shell.execute_reply":"2024-09-23T22:57:28.533944Z"},"trusted":true},"execution_count":null,"outputs":[]}]}