{"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"},{"sourceId":9231709,"sourceType":"datasetVersion","datasetId":5583775}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from sklearn.model_selection import KFold\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\nfrom transformers import get_cosine_schedule_with_warmup\nimport albumentations as A\nimport pydicom\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom glob import glob  # Corrected import\nfrom tqdm import tqdm\n\n# Directories and configurations\nrd = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\nOUTPUT_DIR = '/kaggle/input/rsna24-results'\ndata_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\n\n# Constants\nIMG_SIZE = (512, 512)\nIN_CHANS = 42\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\nMODEL_NAME = \"edgenext_base.in21k_ft_in1k\"\nBATCH_SIZE = 8  # Increased batch size\nN_WORKERS = os.cpu_count()\nUSE_AMP = True\n\n# 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()\n\n# Conditions and Levels\nCONDITIONS = ['spinal_canal_stenosis', 'left_neural_foraminal_narrowing', 'right_neural_foraminal_narrowing',\n              'left_subarticular_stenosis', 'right_subarticular_stenosis']\nLEVELS = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\n# Augmentation Strategy (Train & Test)\ntransforms_test = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.RandomRotate90(),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.2, rotate_limit=15, p=0.5),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, fill_value=0, p=0.5),  # Replacing Cutout\n    A.Normalize(mean=0.5, std=0.5)\n])\n\n# Dataset Definition\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(f'{data_dir}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm'))  # Corrected glob usage\n            img_paths.extend(paths)\n        return img_paths\n\n    def read_dcm_image(self, src_path):\n        try:\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        except Exception as e:\n            print(f\"Error reading DICOM image at {src_path}: {e}\")\n            return np.zeros(IMG_SIZE, dtype=np.uint8)\n\n    def load_series_images(self, study_id, series_desc):\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 image at index {idx} for {series_desc}, study {study_id}: {e}')\n                \n        return images\n\n    def __getitem__(self, idx):\n        study_id = self.study_ids[idx]\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], IN_CHANS), dtype=np.uint8)\n\n        # Load images for each series\n        x[..., :14] = self.load_series_images(study_id, 'Sagittal T1')\n        x[..., 14:28] = self.load_series_images(study_id, 'Sagittal T2/STIR')\n        x[..., 28:] = self.load_series_images(study_id, 'Axial T2')\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# Dataset and DataLoader\ntest_ds = RSNA24TestDataset(df, study_ids, transform=transforms_test)\ntest_dl = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=N_WORKERS, pin_memory=True, drop_last=False)\n\n# Model Definition\nclass RSNA24Model(nn.Module):\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        return self.model(x)\n\n# Load Models\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\nmodels = []\nfor cp in CKPT_PATHS:\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=False)\n    model.load_state_dict(torch.load(cp, map_location=device))\n    model.eval()\n    model.half()\n    model.to(device)\n    models.append(model)\n\n# Autocast for mixed precision\nautocast = torch.amp.autocast(device_type=device, enabled=USE_AMP, dtype=torch.half)\n\n# Predictions and Submission\ny_preds = []\nrow_names = []\n\nwith torch.no_grad():\n    for x_batch, study_ids in tqdm(test_dl):\n        x_batch = x_batch.to(device)\n        batch_preds = np.zeros((x_batch.size(0), N_LABELS, 3))\n\n        with autocast:\n            for model in models:\n                y_batch = model(x_batch)  # Get predictions for the entire batch\n                for i in range(x_batch.size(0)):\n                    for col in range(N_LABELS):\n                        pred = y_batch[i, col * 3: (col + 1) * 3].float().softmax(dim=0).cpu().numpy()\n                        batch_preds[i, col] += pred / len(models)\n\n        # Generate row names for each study in the batch\n        for i, study_id in enumerate(study_ids):\n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(f'{study_id}_{cond}_{level}')\n            y_preds.append(batch_preds[i])\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-10T10:30:01.291539Z","iopub.execute_input":"2024-09-10T10:30:01.291819Z","iopub.status.idle":"2024-09-10T10:30:37.277351Z","shell.execute_reply.started":"2024-09-10T10:30:01.291787Z","shell.execute_reply":"2024-09-10T10:30:37.276205Z"},"trusted":true},"execution_count":null,"outputs":[]}]}