{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import warnings \nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T11:57:54.663628Z","iopub.execute_input":"2025-07-04T11:57:54.663893Z","iopub.status.idle":"2025-07-04T11:57:54.667536Z","shell.execute_reply.started":"2025-07-04T11:57:54.663873Z","shell.execute_reply":"2025-07-04T11:57:54.666858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport matplotlib.pyplot as plt \nimport torch.nn as nn \nimport torch \nfrom  torch.utils.data import Dataset, DataLoader \nfrom torchvision.transforms import Resize\nimport pydicom\nimport glob\nimport os\nimport cv2\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nimport torchvision.transforms as T\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torchvision.transforms as T","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T04:41:45.086389Z","iopub.execute_input":"2025-07-04T04:41:45.087030Z","iopub.status.idle":"2025-07-04T04:41:45.091554Z","shell.execute_reply.started":"2025-07-04T04:41:45.087003Z","shell.execute_reply":"2025-07-04T04:41:45.090852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification'\ntrain_folder = os.path.join(BASE_PATH,'train')\ntest_folder = os.path.join(BASE_PATH,'test')\ntrain_labels = pd.read_csv(os.path.join(BASE_PATH,'train_labels.csv'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T04:41:55.864083Z","iopub.execute_input":"2025-07-04T04:41:55.864358Z","iopub.status.idle":"2025-07-04T04:41:55.904422Z","shell.execute_reply.started":"2025-07-04T04:41:55.864339Z","shell.execute_reply":"2025-07-04T04:41:55.903858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"required_slices=50","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T04:41:56.944342Z","iopub.execute_input":"2025-07-04T04:41:56.944629Z","iopub.status.idle":"2025-07-04T04:41:56.948161Z","shell.execute_reply.started":"2025-07-04T04:41:56.944609Z","shell.execute_reply":"2025-07-04T04:41:56.947450Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels['BraTS21ID'] = train_labels['BraTS21ID'].apply(lambda x: str(x).zfill(5))\ntrain_dict = train_labels.set_index('BraTS21ID')['MGMT_value'].to_dict()\ntrain_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T04:41:57.511236Z","iopub.execute_input":"2025-07-04T04:41:57.511741Z","iopub.status.idle":"2025-07-04T04:41:57.568410Z","shell.execute_reply.started":"2025-07-04T04:41:57.511718Z","shell.execute_reply":"2025-07-04T04:41:57.567694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patient_ids = [ids for ids in list(train_dict.keys()) if ids not in ['00109','00123','00709'] ]\nlen(patient_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T04:41:58.641499Z","iopub.execute_input":"2025-07-04T04:41:58.642197Z","iopub.status.idle":"2025-07-04T04:41:58.646878Z","shell.execute_reply.started":"2025-07-04T04:41:58.642172Z","shell.execute_reply":"2025-07-04T04:41:58.646269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BrDataset(Dataset):\n    def __init__(self, base_path, patient_ids, label_dict, resize_size=(224, 224), required_slices=32, is_train=True):\n        self.base_path = base_path\n        self.patient_ids = patient_ids\n        self.label_dict = label_dict\n        self.modalities = ['FLAIR', 'T1w','T1wCE','T2w']\n        self.required_slices = required_slices\n        self.is_train = is_train\n\n        # Albumentations transform\n        self.transform = A.Compose([\n            A.Resize(height=resize_size[0], width=resize_size[1]),\n            A.HorizontalFlip(p=0.5),\n            A.RandomBrightnessContrast(p=0.1),\n            A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.05, rotate_limit=10, p=0.5),\n            # A.Normalize(mean=0.5, std=0.5),\n        ])\n\n    def load_slices(self, case_id, modality):\n        paths = sorted(\n            glob.glob(os.path.join(self.base_path, 'train', case_id, modality, '*.dcm')),\n            key=lambda x: int(pydicom.dcmread(x).InstanceNumber)\n        )\n\n        slices = []\n        for path in paths:\n            dcm = pydicom.dcmread(path)\n            img = dcm.pixel_array.astype('float32')\n            if np.sum(img) > 0:\n                slices.append(img)\n\n        num_slices = len(slices)\n        if num_slices == 0:\n            return torch.zeros((self.required_slices, 224, 224), dtype=torch.float32)\n\n        # Uniform + Center Sampling\n        if num_slices >= self.required_slices:\n            center = num_slices // 2\n            start = max(0, center - self.required_slices // 2)\n            end = start + self.required_slices\n            if end > num_slices:\n                end = num_slices\n                start = end - self.required_slices\n            indices = np.linspace(start, end - 1, self.required_slices).astype(int)\n            slices = [slices[i] for i in indices]\n        else:\n            pad = self.required_slices - num_slices\n            pad_start = pad // 2\n            pad_end = pad - pad_start\n            h, w = slices[0].shape\n            zero = np.zeros((h, w), dtype=np.float32)\n            slices = [zero] * pad_start + slices + [zero] * pad_end\n\n        volume = []\n        for s in slices:\n            aug = self.transform(image=s)\n            img_tensor = torch.tensor(aug[\"image\"])\n            volume.append(img_tensor)\n\n        volume = torch.stack(volume)\n        return volume\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def __getitem__(self, idx):\n        case_id = str(self.patient_ids[idx]).zfill(5)\n        vol = []\n        for modality in self.modalities:\n            volume = self.load_slices(case_id, modality)\n            vol.append(volume)\n        vol = torch.stack(vol)  # shape: [2, required_slices, H, W]\n\n        label = torch.tensor(self.label_dict[case_id], dtype=torch.long)\n        return vol, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:35:29.663435Z","iopub.execute_input":"2025-07-04T06:35:29.664054Z","iopub.status.idle":"2025-07-04T06:35:29.674518Z","shell.execute_reply.started":"2025-07-04T06:35:29.664029Z","shell.execute_reply":"2025-07-04T06:35:29.673866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = BrDataset(BASE_PATH, patient_ids, train_dict)\nvol,label = dataset[0]\nprint(\"vol shape:\", vol.shape)  # Expect: torch.Size([2, 32, 224, 224])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:35:32.803072Z","iopub.execute_input":"2025-07-04T06:35:32.803595Z","iopub.status.idle":"2025-07-04T06:35:47.245353Z","shell.execute_reply.started":"2025-07-04T06:35:32.803527Z","shell.execute_reply":"2025-07-04T06:35:47.244568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport math\n\ndataset = BrDataset(BASE_PATH, patient_ids, train_dict)\nvol, label = dataset[10]             # vol shape: [2, N, H, W]\nflair_volume = vol[2]                # shape: [N, H, W]\nnum_slices = flair_volume.shape[0]\n\ncols = 6\nrows = math.ceil(num_slices / cols)\n\n# Plot the slices\nfig, axes = plt.subplots(rows, cols, figsize=(16, rows * 2))\nfig.suptitle(f\"FLAIR modality - {num_slices} slices\", fontsize=16)\n\nfor i in range(rows * cols):\n    r, c = divmod(i, cols)\n    ax = axes[r, c] if rows > 1 else axes[c]\n\n    if i < num_slices:\n        slice_2d = flair_volume[i].numpy()\n        ax.imshow(slice_2d, cmap='gray')\n        ax.set_title(f\"Slice {i}\")\n    ax.axis(\"off\")\n\nplt.tight_layout()\nplt.subplots_adjust(top=0.92)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:36:07.720209Z","iopub.execute_input":"2025-07-04T06:36:07.720835Z","iopub.status.idle":"2025-07-04T06:36:12.445384Z","shell.execute_reply.started":"2025-07-04T06:36:07.720812Z","shell.execute_reply":"2025-07-04T06:36:12.444680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass AttentionModule(nn.Module):\n    def __init__(self, input_dim, attention_dim):\n        super().__init__()\n        self.attention = nn.Sequential(\n            nn.Linear(input_dim, attention_dim),\n            nn.Tanh(),\n            nn.Linear(attention_dim, 1)\n        )\n\n    def forward(self, features):  # [B, T, F]\n        scores = self.attention(features)         # [B, T, 1]\n        weights = torch.softmax(scores, dim=1)    # attention weights\n        context = torch.sum(weights * features, dim=1)  # [B, F]\n        return context\n\nclass TumorClassifier3DConvLSTM(nn.Module):\n    def __init__(self, hidden_dim=128, attention_dim=64):\n        super().__init__()\n\n        self.conv3d = nn.Sequential(\n            nn.Conv3d(in_channels=4, out_channels=16, kernel_size=3, padding=1),  # 4 modalities\n            nn.BatchNorm3d(16),\n            nn.ReLU(),\n            nn.MaxPool3d(2),  # → [B, 16, T/2, H/2, W/2]\n            nn.Dropout3d(0.1),\n\n            nn.Conv3d(16, 32, kernel_size=3, padding=1),\n            nn.BatchNorm3d(32),\n            nn.ReLU(),\n            nn.MaxPool3d(2),  # → [B, 32, T/4, H/4, W/4]\n            nn.Dropout3d(0.1),\n        )\n\n        # LSTM will be initialized dynamically\n        self.lstm = nn.LSTM(\n            input_size=1,  # placeholder\n            hidden_size=hidden_dim,\n            num_layers=1,\n            batch_first=True,\n            bidirectional=True\n        )\n\n        self.attention = AttentionModule(hidden_dim * 2, attention_dim)\n\n        self.classifier = nn.Sequential(\n            nn.Linear(hidden_dim * 2, 64),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(64, 2)\n        )\n\n    def forward(self, x):\n        # x: [B, 4, T, H, W]\n        x = self.conv3d(x)  # → [B, 32, T/4, H/4, W/4]\n        b, c, t, h, w = x.shape\n\n        x = x.permute(0, 2, 1, 3, 4)       # → [B, T, C, H, W]\n        x = x.contiguous().view(b, t, -1)  # → [B, T, C*H*W]\n\n        if self.lstm.input_size != x.shape[-1]:\n            self.lstm = nn.LSTM(\n                input_size=x.shape[-1],\n                hidden_size=self.lstm.hidden_size,\n                num_layers=self.lstm.num_layers,\n                batch_first=True,\n                bidirectional=True\n            ).to(x.device)\n\n        lstm_out, _ = self.lstm(x)           # [B, T, 2*hidden_dim]\n        context = self.attention(lstm_out)   # [B, 2*hidden_dim]\n        out = self.classifier(context)       # [B, 2]\n        return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:38:02.514969Z","iopub.execute_input":"2025-07-04T06:38:02.515628Z","iopub.status.idle":"2025-07-04T06:38:02.525432Z","shell.execute_reply.started":"2025-07-04T06:38:02.515604Z","shell.execute_reply":"2025-07-04T06:38:02.524703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\nskf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n\nfor fold, (train_idx, val_idx) in enumerate(skf.split(train_labels['BraTS21ID'], train_labels['MGMT_value'])):\n    train_ids = train_labels.iloc[train_idx]['BraTS21ID'].astype(str).tolist()\n    val_ids = train_labels.iloc[val_idx]['BraTS21ID'].astype(str).tolist()\n    break  # Only using first fold for now\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:38:04.100523Z","iopub.execute_input":"2025-07-04T06:38:04.101037Z","iopub.status.idle":"2025-07-04T06:38:04.108496Z","shell.execute_reply.started":"2025-07-04T06:38:04.101012Z","shell.execute_reply":"2025-07-04T06:38:04.107818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_ids),len(val_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:38:06.352139Z","iopub.execute_input":"2025-07-04T06:38:06.352702Z","iopub.status.idle":"2025-07-04T06:38:06.357224Z","shell.execute_reply.started":"2025-07-04T06:38:06.352676Z","shell.execute_reply":"2025-07-04T06:38:06.356655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom torch.utils.data import DataLoader\nimport torch.nn.functional as F\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nfrom tqdm import tqdm\nimport torch\n\n# =================== SETTINGS ===================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nBATCH_SIZE = 2\nEPOCHS = 10\nLR = 1e-4\nPATIENCE = 3\nBEST_MODEL_PATH = \"/kaggle/working/best_model.pt\"\n\n# =================== DATALOADERS ===================\ntrain_dataset = BrDataset(BASE_PATH, train_ids, train_dict, resize_size=(224, 224), required_slices=32, is_train=True)\nval_dataset = BrDataset(BASE_PATH, val_ids, train_dict, resize_size=(224, 224), required_slices=32, is_train=False)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\n# =================== MODEL + LOSS + OPTIM ===================\nmodel = TumorClassifier3DConvLSTM().to(device)\ncriterion = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\n# =================== TRAINING LOOP WITH EARLY STOPPING ===================\ndef train_and_validate(model, train_loader, val_loader, criterion, optimizer, device, epochs=10):\n    best_auc = 0\n    patience_counter = 0\n\n    for epoch in range(epochs):\n        print(f\"\\n--- Epoch {epoch+1}/{epochs} ---\")\n\n        # -------- TRAIN --------\n        model.train()\n        total_loss = 0\n\n        for x, y in tqdm(train_loader):\n            x, y = x.to(device), y.to(device).long()\n\n            optimizer.zero_grad()\n            out = model(x)\n            loss = criterion(out, y)\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n\n        avg_loss = total_loss / len(train_loader)\n        print(f\"Train Loss: {avg_loss:.4f}\")\n\n        # -------- VALIDATION --------\n        model.eval()\n        val_preds, val_targets = [], []\n\n        with torch.no_grad():\n            for x, y in val_loader:\n                x, y = x.to(device), y.to(device).long()\n                out = model(x)\n                prob = F.softmax(out, dim=1)[:, 1]\n\n                val_preds.extend(prob.cpu().numpy())\n                val_targets.extend(y.cpu().numpy())\n\n        val_auc = roc_auc_score(val_targets, val_preds)\n        val_preds_cls = (torch.tensor(val_preds) > 0.5).int()\n        val_acc = accuracy_score(val_targets, val_preds_cls)\n\n        print(f\"Val AUC: {val_auc:.4f} | Val Accuracy: {val_acc:.4f}\")\n\n        # -------- EARLY STOPPING & CHECKPOINT --------\n        if val_auc > best_auc:\n            print(f\"AUC improved from {best_auc:.4f} → {val_auc:.4f}. Saving model.\")\n            best_auc = val_auc\n            patience_counter = 0\n            torch.save(model.state_dict(), BEST_MODEL_PATH)\n        else:\n            patience_counter += 1\n            print(f\"No improvement. Patience: {patience_counter}/{PATIENCE}\")\n\n        if patience_counter >= PATIENCE:\n            print(\"Early stopping triggered.\")\n            break\n\n# =================== RUN ===================\ntrain_and_validate(model, train_loader, val_loader, criterion, optimizer, device, epochs=EPOCHS)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-04T06:38:07.739438Z","iopub.execute_input":"2025-07-04T06:38:07.740166Z","iopub.status.idle":"2025-07-04T11:40:33.841084Z","shell.execute_reply.started":"2025-07-04T06:38:07.740141Z","shell.execute_reply":"2025-07-04T11:40:33.830685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}