{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":129601,"databundleVersionId":15542776,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14695456,"sourceType":"datasetVersion","datasetId":9387785}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":30601.931586,"end_time":"2026-02-02T22:13:15.439804","environment_variables":{},"exception":true,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-02-02T13:43:13.508218","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Brain Tumor Segmentation using SwinUNETR\n\n## Project Overview\nThis notebook implements a deep learning pipeline to segment brain tumors from MRI scans using the BraTS dataset. We utilize a SwinUNETR architecture to predict segmentation masks from multi-modal MRI inputs.\n\n## Objectives\n1.  **Data Loading**: Load 3D MRI volumes (NIfTI format).\n2.  **Preprocessing**: Normalize and format data for the model.\n3.  **Model**: Implement a SwinUNETR architecture in PyTorch.\n4.  **Training**: Train the model using Dice Loss.\n5.  **Visualization**: Visualize inputs, ground truth, and predictions.\n","metadata":{"papermill":{"duration":0.00326,"end_time":"2026-02-02T13:43:16.975118","exception":false,"start_time":"2026-02-02T13:43:16.971858","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## 1. Setup and Imports\nInstall necessary packages if not already present.\n","metadata":{"papermill":{"duration":0.00249,"end_time":"2026-02-02T13:43:16.979993","exception":false,"start_time":"2026-02-02T13:43:16.977503","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install monai nibabel matplotlib torch torchvision tqdm scikit-learn comet_ml -q","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:34:32.236553Z","iopub.status.busy":"2026-02-03T17:34:32.236323Z","iopub.status.idle":"2026-02-03T17:34:42.876264Z","shell.execute_reply":"2026-02-03T17:34:42.875521Z","shell.execute_reply.started":"2026-02-03T17:34:32.236534Z"},"papermill":{"duration":10.463384,"end_time":"2026-02-02T13:43:27.445557","exception":false,"start_time":"2026-02-02T13:43:16.982173","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import comet_ml\nfrom comet_ml import Experiment\nimport os\nimport glob\nimport gc\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom monai import transforms\nfrom monai import data\n\nfrom monai.networks.nets import SwinUNETR, DynUNet\nfrom monai.losses import DiceLoss\nfrom monai.inferers import sliding_window_inference\n\n# Check for GPU\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nif device.type == 'cuda':\n    torch.backends.cudnn.benchmark = True\nprint(f\"Using device: {device}\")\n\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom monai.metrics import DiceMetric\nfrom monai.data import decollate_batch\nfrom monai.transforms import AsDiscrete, Compose, EnsureType, Activations\n","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:34:42.878404Z","iopub.status.busy":"2026-02-03T17:34:42.878147Z","iopub.status.idle":"2026-02-03T17:35:30.182842Z","shell.execute_reply":"2026-02-03T17:35:30.182041Z","shell.execute_reply.started":"2026-02-03T17:34:42.878377Z"},"papermill":{"duration":51.326389,"end_time":"2026-02-02T13:44:18.774693","exception":false,"start_time":"2026-02-02T13:43:27.448304","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Data Configuration","metadata":{"papermill":{"duration":0.002537,"end_time":"2026-02-02T13:44:18.779963","exception":false,"start_time":"2026-02-02T13:44:18.777426","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# === CONFIGURATION ===\n# Initialize Comet ML Experiment\nexperiment = Experiment(\n    api_key = \"VDWJEDjieOQbqDskMAw7Cr1av\",\n    project_name=\"brain-tumor-segmentation\",\n    auto_metric_logging=True,\n    auto_param_logging=True,\n    auto_histogram_weight_logging=True,\n    auto_histogram_gradient_logging=True,\n    auto_histogram_activation_logging=True,\n)\n\nDATASET_ROOT = \"/kaggle/input/brain-tumor-segmentation-hackathon\"\n\n# Parameters\nIMG_SIZE = 128\nBATCH_SIZE = 1\nLEARNING_RATE = 5e-5\nNUM_EPOCHS = 100\nVALIDATION_SPLIT = 0.2\n\n# Log parameters\nparameters = {\n    \"IMG_SIZE\":IMG_SIZE,\n    \"batch_size\": BATCH_SIZE,\n    \"learning_rate\": LEARNING_RATE,\n    \"epochs\": NUM_EPOCHS,\n    \"validation_split\": VALIDATION_SPLIT,\n}\nexperiment.log_parameters(parameters)\n\ntest_index = 0      # Index for visualization","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:30.187802Z","iopub.status.busy":"2026-02-03T17:35:30.187514Z","iopub.status.idle":"2026-02-03T17:35:35.779938Z","shell.execute_reply":"2026-02-03T17:35:35.778928Z","shell.execute_reply.started":"2026-02-03T17:35:30.187778Z"},"papermill":{"duration":5.171803,"end_time":"2026-02-02T13:44:23.954362","exception":false,"start_time":"2026-02-02T13:44:18.782559","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Dataset Implementation\nWe define a custom `BraTSDataset` class.\n-   **Inputs**: Reads 4 modalites: FLAIR, T1, T1ce, T2.\n-   **Labels**: \n    -   Label 1: Necrotic/core (NCR)\n    -   Label 2: Edema (ED)\n    -   Label 4: Enhancing Tumor (ET)\n    -   *Mapped to 0, 1, 2, 3 for training*.\n","metadata":{"papermill":{"duration":0.003456,"end_time":"2026-02-02T13:44:23.961135","exception":false,"start_time":"2026-02-02T13:44:23.957679","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class BraTSDataset(Dataset):\n    def __init__(self, root_dir, transform=None, img_size=128, mode='train', test_dir=None):\n        self.root_dir = root_dir\n        self.transform = transform\n        self.img_size = img_size\n        self.mode = mode\n        \n        if mode == 'test':\n            self.case_folders = sorted(glob.glob(os.path.join(test_dir, \"BraTS2021_*\")))\n            print(f\"Found {len(self.case_folders)} cases for test\")\n        else:\n            # Get all case folders\n            self.case_folders = sorted(glob.glob(os.path.join(root_dir, \"BraTS2021_*\")))\n                    \n            # Simple train/val split based on index\n            split_idx = int(len(self.case_folders) * (1 - VALIDATION_SPLIT))\n            if mode == 'train':\n                self.case_folders = self.case_folders[:split_idx]\n            else:\n                self.case_folders = self.case_folders[split_idx:]\n            \n            print(f\"Found {len(self.case_folders)} cases for {mode}\")\n        \n        # Pre-load file paths\n        self.data_list = []\n        print(f\"Pre-loading file paths for {mode}...\")\n        for case_path in tqdm(self.case_folders):\n            case_id = os.path.basename(case_path)\n            try:\n                flair_path = self._get_file_path(case_path, case_id, \"flair\")\n                t1_path = self._get_file_path(case_path, case_id, \"t1\")\n                t1ce_path = self._get_file_path(case_path, case_id, \"t1ce\")\n                t2_path = self._get_file_path(case_path, case_id, \"t2\")\n                \n                item = {\n                    \"image\": [flair_path, t1_path, t1ce_path, t2_path],\n                    \"id\": case_id\n                }\n                \n                if mode != 'test':\n                    seg_path = self._get_file_path(case_path, case_id, \"seg\")\n                    item[\"label\"] = seg_path\n                \n                self.data_list.append(item)\n                \n            except Exception as e:\n                print(f\"Skipping case {case_id}: {e}\")\n    \n    def __len__(self):\n        return len(self.data_list)\n\n    def _get_file_path(self, case_path, case_id, modality):\n        # Helper to find the .nii file for a given modality.\n        # Structure: Case/Case_modality.nii/file.nii\n        \n        folder_name = f\"{case_id}_{modality}.nii\"\n        folder_path = os.path.join(case_path, folder_name)\n        \n        # 1. Try to find the file in the specific folder (Nested structure)\n        if os.path.isdir(folder_path):\n            files = glob.glob(os.path.join(folder_path, \"*.nii\"))\n            if files:\n                if len(files) > 1:\n                    print(f\"Multiple files found in {folder_path}: {len(files)}\")\n                return files[0]\n                \n        # 2. Fallback: Check if the folder_path itself is actually the file\n        if os.path.isfile(folder_path):\n            return folder_path\n\n        # 3. Final Fallback: If not found, search the parent case_path for all available .nii files\n        available_files = glob.glob(os.path.join(case_path, \"**\", \"*.nii\"), recursive=True)\n      \n        raise FileNotFoundError(f\"File for modality '{modality}' not found in {case_path}\")\n\n    def __getitem__(self, idx):\n        data = self.data_list[idx]\n        if self.transform:\n            data = self.transform(data)\n        return data\n","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:35.782080Z","iopub.status.busy":"2026-02-03T17:35:35.781845Z","iopub.status.idle":"2026-02-03T17:35:35.800364Z","shell.execute_reply":"2026-02-03T17:35:35.799453Z","shell.execute_reply.started":"2026-02-03T17:35:35.782059Z"},"papermill":{"duration":0.016116,"end_time":"2026-02-02T13:44:23.980120","exception":false,"start_time":"2026-02-02T13:44:23.964004","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_loader(batch_size, data_dir, roi):\n    # Initialize datasets\n    train_files = BraTSDataset(data_dir, mode='train')\n    validation_files = BraTSDataset(data_dir, mode='val')\n    train_transform = transforms.Compose(\n        [\n            transforms.LoadImaged(keys=[\"image\", \"label\"]),\n            transforms.EnsureChannelFirstd(keys=[\"image\"]),\n            transforms.ConvertToMultiChannelBasedOnBratsClassesd(keys=\"label\"),\n            transforms.CropForegroundd(\n                keys=[\"image\", \"label\"],\n                source_key=\"image\",\n                k_divisible=[roi[0], roi[1], roi[2]],\n                allow_smaller=True,\n            ),\n            transforms.RandSpatialCropd(\n                keys=[\"image\", \"label\"],\n                roi_size=[roi[0], roi[1], roi[2]],\n                random_size=False,\n            ),\n            transforms.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=0),\n            transforms.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=1),\n            transforms.RandFlipd(keys=[\"image\", \"label\"], prob=0.5, spatial_axis=2),\n            transforms.NormalizeIntensityd(keys=\"image\", nonzero=True, channel_wise=True),\n            transforms.RandScaleIntensityd(keys=\"image\", factors=0.1, prob=1.0),\n            transforms.RandShiftIntensityd(keys=\"image\", offsets=0.1, prob=1.0),\n        ]\n    )\n    val_transform = transforms.Compose( \n        [\n            transforms.LoadImaged(keys=[\"image\", \"label\"]),\n            transforms.EnsureChannelFirstd(keys=[\"image\"]),\n            transforms.ConvertToMultiChannelBasedOnBratsClassesd(keys=\"label\"),\n            transforms.NormalizeIntensityd(keys=\"image\", nonzero=True, channel_wise=True),\n        ]\n    )\n\n    train_ds = data.Dataset(data=train_files, transform=train_transform)\n\n    train_loader = data.DataLoader(\n        train_ds,\n        batch_size=batch_size,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=True,\n    )\n    val_ds = data.Dataset(data=validation_files, transform=val_transform)\n    val_loader = data.DataLoader(\n        val_ds,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True,\n    )\n\n    return train_loader, val_loader","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:35.803299Z","iopub.status.busy":"2026-02-03T17:35:35.803026Z","iopub.status.idle":"2026-02-03T17:35:35.835989Z","shell.execute_reply":"2026-02-03T17:35:35.835366Z","shell.execute_reply.started":"2026-02-03T17:35:35.803265Z"},"papermill":{"duration":0.012411,"end_time":"2026-02-02T13:44:23.995208","exception":false,"start_time":"2026-02-02T13:44:23.982797","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training\n","metadata":{"papermill":{"duration":0.002667,"end_time":"2026-02-02T13:44:24.000614","exception":false,"start_time":"2026-02-02T13:44:23.997947","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Initialize Model\nmodel = SwinUNETR(\n    in_channels=4,\n    out_channels=3,\n    feature_size=48,\n    drop_rate=0.0,\n    attn_drop_rate=0.0,\n    dropout_path_rate=0.0,\n    use_checkpoint=True,\n).to(device)\ncriterion = DiceLoss(sigmoid=True)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=1e-5)\nscheduler = CosineAnnealingLR(optimizer, T_max=NUM_EPOCHS)\n\nprint(\"Model and Loaders initialized.\")","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:35.837111Z","iopub.status.busy":"2026-02-03T17:35:35.836910Z","iopub.status.idle":"2026-02-03T17:35:36.639372Z","shell.execute_reply":"2026-02-03T17:35:36.638593Z","shell.execute_reply.started":"2026-02-03T17:35:35.837092Z"},"papermill":{"duration":0.786717,"end_time":"2026-02-02T13:44:24.790040","exception":false,"start_time":"2026-02-02T13:44:24.003323","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Download Official BraTS 21 Pretrained Weights (Fine-tuned)\nimport os\nimport torch\nimport requests\nimport zipfile\nimport numpy\nfrom tqdm import tqdm\n\n# URL for Fold 0 (matches feature_size=48)\nmodel_url = \"https://github.com/Project-MONAI/MONAI-extra-test-data/releases/download/0.8.1/fold0_f48_ep300_4gpu_dice0_8854.zip\"\nzip_path = \"model_brats21.zip\"\nextract_path = \"model_brats21\"\n\nif not os.path.exists(zip_path) and not os.path.exists(extract_path):\n    print(f\"Downloading {zip_path}...\")\n    response = requests.get(model_url, stream=True)\n    total_size_in_bytes = int(response.headers.get('content-length', 0))\n    block_size = 1024 # 1 Kibibyte\n    progress_bar = tqdm(total=total_size_in_bytes, unit='iB', unit_scale=True)\n    \n    with open(zip_path, 'wb') as f:\n        for chunk in response.iter_content(block_size):\n            progress_bar.update(len(chunk))\n            f.write(chunk)\n    progress_bar.close()\n    \n    # Unzip\n    print(\"Unzipping model...\")\n    with zipfile.ZipFile(zip_path, 'r') as zip_ref:\n        zip_ref.extractall(extract_path)\n    print(\"Unzip complete.\")\n\n# Find the model file (it might have a long name inside the zip)\nmodel_file = None\nif os.path.exists(extract_path):\n    for root, dirs, files in os.walk(extract_path):\n        for file in files:\n            if file.endswith(\".pt\") or file.endswith(\".pth\"):\n                model_file = os.path.join(root, file)\n                break\n\nif model_file:\n    print(f\"Loading pretrained weights from {model_file}...\")\n    try:\n        # Fix for WeightsUnpickler error with numpy scalars\n        if hasattr(torch.serialization, 'add_safe_globals'):\n            torch.serialization.add_safe_globals([numpy.core.multiarray.scalar])\n\n        # Load state dict\n        checkpoint = torch.load(model_file, map_location=device, weights_only=False)\n        if \"state_dict\" in checkpoint:\n            pretrained_state_dict = checkpoint[\"state_dict\"]\n        else:\n            pretrained_state_dict = checkpoint\n        \n        # Load full model (expecting all keys to match)\n        model.load_state_dict(pretrained_state_dict, strict=True)\n        \n        print(\"Successfully loaded BraTS 21 fine-tuned weights (Encoder + Head).\")\n        print(\"Model is now ready to verify or run inference.\")\n        \n    except Exception as e:\n        print(f\"Error loading pretrained weights: {e}\")\n        # Fallback to loose loading just in case\n        print(\"Attempting loose loading...\")\n        model.load_state_dict(pretrained_state_dict, strict=False)\n\nelse:\n    print(\"Model file not found after extraction.\")\n","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:36.640845Z","iopub.status.busy":"2026-02-03T17:35:36.640404Z","iopub.status.idle":"2026-02-03T17:35:46.569649Z","shell.execute_reply":"2026-02-03T17:35:46.568642Z","shell.execute_reply.started":"2026-02-03T17:35:36.640821Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Enhanced Training function with Comet ML Logging and DiceMetric\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs, patience=5, val_interval=1, save_path='best_model.pth'):\n    best_val_dice = -1.0\n    epochs_no_improve = 0\n    scaler = torch.amp.GradScaler('cuda') # Initialize GradScaler for AMP\n    \n    # Initialize Metrics\n    dice_metric = DiceMetric(include_background=True, reduction=\"mean\")\n    dice_metric_batch = DiceMetric(include_background=True, reduction=\"mean_batch\")\n    \n    # Post-processing transforms for validation\n    post_trans = Compose([Activations(sigmoid=True), AsDiscrete(threshold=0.5)])\n    \n    with experiment.train():\n        step = 0\n        for epoch in range(num_epochs):\n            # === TRAINING ===\n            model.train()\n            train_loss = 0\n            loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Train]\")\n            \n            for batch_data in loop:\n                images = batch_data[\"image\"].to(device)\n                masks = batch_data[\"label\"].to(device)\n                \n                # Forward with AMP\n                with torch.amp.autocast('cuda'):\n                    outputs = model(images)\n                    loss = criterion(outputs, masks)\n                \n                # Backward\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n                \n                train_loss += loss.item()\n                loop.set_postfix(loss=loss.item())\n                \n                # Log batch loss\n                experiment.log_metric(\"batch_train_loss\", loss.item(), step=step)\n                step += 1\n                \n            avg_train_loss = train_loss / len(train_loader)\n            experiment.log_metric(\"epoch_train_loss\", avg_train_loss, step=epoch)\n            \n            # Optimization: Clear memory\n            del images, masks, outputs, loss\n            gc.collect()\n            torch.cuda.empty_cache()\n            \n            # === VALIDATION ===\n            if (epoch + 1) % val_interval == 0:\n                model.eval()\n                val_loss = 0\n                with torch.no_grad():\n                    with experiment.validate():\n                        for batch_data in val_loader:\n                            images = batch_data[\"image\"].to(device)\n                            masks = batch_data[\"label\"].to(device)\n                            \n                            with torch.amp.autocast('cuda'):\n                                # Use overlap=0.25 and sw_batch_size=4 for faster validation\n                                val_outputs = sliding_window_inference(images, (IMG_SIZE, IMG_SIZE, IMG_SIZE), 3, model, overlap=0.25)\n                                loss = criterion(val_outputs, masks)\n                            \n                            val_loss += loss.item()\n                            \n                            # Compute Dice Metric\n                            val_outputs = [post_trans(i) for i in decollate_batch(val_outputs)]\n                            val_labels = decollate_batch(masks)\n                            dice_metric(y_pred=val_outputs, y=val_labels)\n                            dice_metric_batch(y_pred=val_outputs, y=val_labels)\n                \n                avg_val_loss = val_loss / len(val_loader)\n                \n                # Aggregate Metrics\n                mean_val_dice = dice_metric.aggregate().item()\n                metric_batch = dice_metric_batch.aggregate()\n                \n                # Metric per class (TC, WT, ET)\n                metric_tc = metric_batch[0].item()\n                metric_wt = metric_batch[1].item()\n                metric_et = metric_batch[2].item()\n                \n                dice_metric.reset()\n                dice_metric_batch.reset()\n                \n                # Log Metrics\n                experiment.log_metric(\"epoch_val_loss\", avg_val_loss, step=epoch)\n                experiment.log_metric(\"val_mean_dice\", mean_val_dice, step=epoch)\n                experiment.log_metric(\"val_dice_TC\", metric_tc, step=epoch)\n                experiment.log_metric(\"val_dice_WT\", metric_wt, step=epoch)\n                experiment.log_metric(\"val_dice_ET\", metric_et, step=epoch)\n                \n                print(f\"Epoch {epoch+1}: Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f} | Val Mean Dice: {mean_val_dice:.4f}\")\n                print(f\"Dice Classes -> TC: {metric_tc:.4f}, WT: {metric_wt:.4f}, ET: {metric_et:.4f}\")\n                \n                # === SAVE BEST MODEL (Based on Dice) ===\n                if mean_val_dice > best_val_dice:\n                    best_val_dice = mean_val_dice\n                    epochs_no_improve = 0\n                    torch.save(model.state_dict(), save_path)\n                    experiment.log_model(\"UNet_Best\", save_path)\n                    print(f\"Validation Dice improved. Model saved to {save_path}\")\n                else:\n                    epochs_no_improve += 1\n                    print(f\"No improvement for {epochs_no_improve} checks.\")\n                    \n                # === EARLY STOPPING ===\n                if epochs_no_improve >= patience:\n                    print(\"Early stopping triggered!\")\n                    break\n            else:\n                print(f\"Epoch {epoch+1}: Train Loss: {avg_train_loss:.4f} | Validation Skipped\")\n            \n            # === SCHEDULER ===\n            scheduler.step()\n                \n            gc.collect()\n            torch.cuda.empty_cache()\n            \ntrain_loader, val_loader = get_loader(BATCH_SIZE, DATASET_ROOT, roi=(IMG_SIZE, IMG_SIZE, IMG_SIZE))\ntrain_model(model, train_loader, val_loader, criterion, optimizer, scheduler, NUM_EPOCHS, patience=3, val_interval=3)","metadata":{"execution":{"iopub.execute_input":"2026-02-03T17:35:46.571166Z","iopub.status.busy":"2026-02-03T17:35:46.570793Z","iopub.status.idle":"2026-02-03T17:36:05.531418Z","shell.execute_reply":"2026-02-03T17:36:05.530156Z","shell.execute_reply.started":"2026-02-03T17:35:46.571138Z"},"papermill":{"duration":30523.63607,"end_time":"2026-02-02T22:13:08.428962","exception":false,"start_time":"2026-02-02T13:44:24.792892","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Visualization\nVisualize a sample prediction.\n","metadata":{"papermill":{"duration":0.575445,"end_time":"2026-02-02T22:13:09.573512","exception":false,"start_time":"2026-02-02T22:13:08.998067","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Visualize and Log to Comet\ndef visualize_prediction(model, dataset, index=0):\n    model.eval()\n    data = dataset[index]\n    image = data['image']\n    mask = data['label']\n    \n    # Prepare input\n    input_tensor = image.unsqueeze(0).to(device)\n    with torch.no_grad():\n        output = sliding_window_inference(input_tensor, (IMG_SIZE, IMG_SIZE, IMG_SIZE), 3, model, overlap=0.5)\n        # Output is (1, 3, D, H, W). Use sigmoid > 0.5 for regions\n        prediction = (torch.sigmoid(output) > 0.5).float().cpu().numpy()[0]\n        \n    # Visualize Middle Slice\n    # Find slice with maximum label area (Enhancing Tumor)\n    if mask.sum() > 0:\n        # Try Enhancing Tumor (channel 2)\n        slice_areas = mask[2].sum(dim=(0, 1))\n        if slice_areas.max() > 0:\n            slice_idx = torch.argmax(slice_areas).item()\n        else:\n            # Any tumor class\n            slice_areas = mask.sum(dim=0).sum(dim=(0, 1))\n            if slice_areas.max() > 0:\n                slice_idx = torch.argmax(slice_areas).item()\n            else:\n                slice_idx = image.shape[-1] // 2\n    else:\n        slice_idx = image.shape[-1] // 2\n    print(f\"Visualizing Slice: {slice_idx}\")\n    \n    fig, ax = plt.subplots(1, 4, figsize=(20, 5))\n    \n    # Show FLAIR channel (channel 0)\n    ax[0].imshow(image[0, :, :, slice_idx].cpu(), cmap='gray')\n    ax[0].set_title(\"Input (FLAIR)\")\n    \n    # Show T1ce channel (channel 2)\n    ax[1].imshow(image[2, :, :, slice_idx].cpu(), cmap='gray')\n    ax[1].set_title(\"Input (T1ce)\")\n\n    # Show Ground Truth (Enhancing Tumor - channel 2)\n    ax[2].imshow(mask[2, :, :, slice_idx].cpu(), cmap='gray')\n    ax[2].set_title(\"Ground Truth (ET)\")\n    \n    # Show Prediction (Enhancing Tumor - channel 2)\n    ax[3].imshow(prediction[2, :, :, slice_idx], cmap='gray')\n    ax[3].set_title(\"Prediction (ET)\")\n    \n    # Save figure to log\n    plt.savefig(\"prediction_sample.png\")\n    experiment.log_image(\"prediction_sample.png\", name=f\"Prediction Sample Index {index}\")\n    \n    plt.show()\n\n# Visualize\nvisualize_prediction(model, val_loader.dataset, index=0)","metadata":{"execution":{"iopub.status.busy":"2026-02-03T17:36:05.532033Z","iopub.status.idle":"2026-02-03T17:36:05.532321Z","shell.execute_reply":"2026-02-03T17:36:05.532212Z","shell.execute_reply.started":"2026-02-03T17:36:05.532196Z"},"papermill":{"duration":1.409115,"end_time":"2026-02-02T22:13:11.761773","exception":true,"start_time":"2026-02-02T22:13:10.352658","status":"failed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# === SUBMISSION PIPELINE ===\ndef rle_encoding(x):\n    '''\n    x: numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns run length as list\n    '''\n    dots = np.where(x.flatten() == 1)[0] # .flatten() gives row-major flattening\n    if len(dots) == 0:\n        return \"\"\n    \n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if (b > prev + 1):\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n        \n    return \" \".join([str(x) for x in run_lengths])\n\ndef run_inference_and_submit(model, test_dir, output_file=\"submission.csv\"):\n    test_transform = transforms.Compose([\n        transforms.LoadImaged(keys=[\"image\"]),\n        transforms.EnsureChannelFirstd(keys=[\"image\"]),\n        transforms.NormalizeIntensityd(keys=\"image\", nonzero=True, channel_wise=True),\n    ])\n    \n    test_dataset = BraTSDataset(root_dir=\"\", transform=test_transform, mode='test', test_dir=test_dir)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    model.eval()\n    \n    results = []\n    \n    with open(output_file, 'w') as f:\n        f.write(\"id,rle\\n\")\n        \n        with torch.no_grad():\n            for batch_data in tqdm(test_loader, desc=\"Inference\"):\n                images = batch_data[\"image\"].to(device)\n                case_id = batch_data[\"id\"][0]\n                \n                # Inference\n                outputs = sliding_window_inference(images, (IMG_SIZE, IMG_SIZE, IMG_SIZE), 3, model, overlap=0.5)\n                \n                # Output channels: 0: TC, 1: WT, 2: ET\n                probs = torch.sigmoid(outputs)\n                preds = (probs > 0.5).float().cpu().numpy()[0] # (3, D, H, W)\n                \n                # Reconstruct Classes for Submission\n                # Class 1 (NCR): TC (0) - ET (2)\n                # Class 2 (ED): WT (1) - TC (0)\n                # Class 4 (ET): ET (2)\n                \n                tc = preds[0]\n                wt = preds[1]\n                et = preds[2]\n                \n                label_1 = (tc - et)\n                label_1[label_1 < 0] = 0 # Safety clip\n                \n                label_2 = (wt - tc)\n                label_2[label_2 < 0] = 0\n                \n                label_4 = et\n                \n                # Encode and Write\n                # 1: NCR\n                rle_1 = rle_encoding(label_1)\n                f.write(f\"{case_id}_1,{rle_1}\\n\")\n                \n                # 2: ED\n                rle_2 = rle_encoding(label_2)\n                f.write(f\"{case_id}_2,{rle_2}\\n\")\n                \n                # 4: ET\n                rle_4 = rle_encoding(label_4)\n                f.write(f\"{case_id}_4,{rle_4}\\n\")\n                \n    print(f\"Submission saved to {output_file}\")\n\n# model_path = \"/kaggle/input/brain-tumor-segmentation/best_model.pth\"\n# model.load_state_dict(torch.load(model_path, map_location=device))\n# print(f\"Loaded model from {model_path}\")\n\n# Run Submission Pipeline\nTEST_DIR = \"/kaggle/input/instant-odc-ai-hackathon/test\"\nrun_inference_and_submit(model, TEST_DIR)","metadata":{"execution":{"iopub.status.busy":"2026-02-03T17:36:05.533783Z","iopub.status.idle":"2026-02-03T17:36:05.534058Z","shell.execute_reply":"2026-02-03T17:36:05.533947Z","shell.execute_reply.started":"2026-02-03T17:36:05.533930Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"experiment.end()","metadata":{"execution":{"iopub.status.busy":"2026-02-03T17:36:05.535776Z","iopub.status.idle":"2026-02-03T17:36:05.536099Z","shell.execute_reply":"2026-02-03T17:36:05.535959Z","shell.execute_reply.started":"2026-02-03T17:36:05.535933Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}