{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"sourceType":"competition"}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==============================================================================\n# Environment Setup for Swin UNETR (Corrected)\n# Run this in the VERY FIRST cell, then RESTART the session.\n# ==============================================================================\nimport os\n\nprint(\"--> Step 1: Uninstalling default Kaggle libraries to avoid conflicts...\")\n!pip uninstall -y torch torchvision torchaudio monai torchtext\n\nprint(\"\\n--> Step 2: Installing a specific, stable version of PyTorch (1.13.1)...\")\n# We only pin the main torch library version.\n!pip install torch==1.13.1+cu117 --extra-index-url https://download.pytorch.org/whl/cu117\n\nprint(\"\\n--> Step 3: Installing a compatible MONAI version (1.1.0)...\")\n# Pip will now automatically find the correct torchvision/torchaudio versions for us.\n!pip install \"monai[nibabel, tqdm]==1.1.0\"\n\nprint(\"\\n--> Step 4: Environment setup is complete.\")\nprint(\"!!! IMPORTANT: You MUST restart the session now for these changes to take effect. !!!\")\n\n# This line will cause the session to crash, forcing a restart.\n#os.kill(os.getpid(), 9)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-10T07:34:34.99484Z","iopub.execute_input":"2025-10-10T07:34:34.995093Z","execution_failed":"2025-10-10T07:36:34.585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Step 2: Full Training and Inference Script with Metric Plotting\n# ==============================================================================\n!pip install -q 'monai[nibabel, tqdm]'\nimport os\nimport glob\nimport gc\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport nibabel as nib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport warnings\nimport cv2\n\n# MONAI Imports\nfrom monai.data import DataLoader, CacheDataset, decollate_batch\nfrom monai.losses import DiceLoss, FocalLoss\nfrom monai.metrics import DiceMetric\nfrom monai.networks.nets import SwinUNETR\nfrom monai.inferers import sliding_window_inference\nfrom monai.transforms import (\n    Compose, LoadImaged, EnsureChannelFirstd, Orientationd, Spacingd,\n    CropForegroundd, RandCropByPosNegLabeld, ToTensord, Lambdad, SpatialPadd,\n    SelectItemsd, AsDiscrete, MapTransform, ScaleIntensityRanged, DivisiblePadd\n)\nfrom monai.utils import set_determinism\nfrom torch.hub import load_state_dict_from_url\n\nwarnings.filterwarnings(\"ignore\", message=\".*unable to generate class balanced samples.*\")\nos.environ['CUDA_LAUNCH_BLOCKING'] = '1'\nset_determinism(seed=42)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Configuration\n# ==============================================================================\nCACHE_RATE = 0.25\nNUM_EPOCHS = 25\nPATCH_SIZE = (64, 64, 64) \nBEST_MODEL_FILENAME = \"best_model.pth\"\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Phase 2: Data Loading and Preparation\n# ==============================================================================\nprint(\"--- Phase 2: Loading and Preparing Data Paths ---\")\ndata_dir = \"/kaggle/input/rsna-intracranial-aneurysm-detection\"\ntrain_df = pd.read_csv(os.path.join(data_dir, \"train.csv\"))\nmodality_map = train_df.set_index('SeriesInstanceUID')['Modality'].to_dict()\nsegmentation_folder = os.path.join(data_dir, \"segmentations\")\nsegmentation_files = glob.glob(os.path.join(segmentation_folder, \"*_cowseg.nii\"))\n\ndata_dicts = []\nfor label_path in segmentation_files:\n    image_path = label_path.replace(\"_cowseg.nii\", \".nii\")\n    series_uid = os.path.basename(image_path).replace(\".nii\", \"\")\n    if os.path.exists(image_path) and series_uid in modality_map:\n        data_dicts.append({\n            \"image\": image_path, \n            \"label\": label_path, \n            \"modality\": modality_map[series_uid]\n        })\n\nprint(f\"Found {len(data_dicts)} matching image-label-modality sets.\")\nsplit_index = int(len(data_dicts) * 0.8)\ntrain_files, val_files = data_dicts[:split_index], data_dicts[split_index:]\nprint(f\"Training samples: {len(train_files)}, Validation samples: {len(val_files)}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Phase 3: Custom MONAI Transforms\n# ==============================================================================\nclass ModalityNormalize(MapTransform):\n    def __call__(self, data):\n        d = dict(data)\n        modality = d.get(\"modality\", \"MRA\")\n        image_key = self.keys[0]\n        image_data = d[image_key]\n        \n        if modality == \"CTA\":\n            a_min, a_max = 0, 400\n        else:\n            pixels = image_data[image_data > 0]\n            if len(pixels) > 0:\n                a_min, a_max = np.percentile(pixels, 1), np.percentile(pixels, 99)\n            else:\n                a_min, a_max = 0, 0\n        \n        if a_max <= a_min:\n            a_min, a_max = image_data.min(), image_data.max()\n            \n        normalizer = ScaleIntensityRanged(\n            keys=image_key, a_min=a_min, a_max=a_max, b_min=0.0, b_max=1.0, clip=True\n        )\n        return normalizer(d)\n\nempty_label_count = 0\nempty_label_uids = []\nclass LogEmptyLabelsd(MapTransform):\n    def __call__(self, data):\n        d = dict(data)\n        if np.sum(d[\"label\"]) == 0:\n            global empty_label_count, empty_label_uids\n            empty_label_count += 1\n            uid = os.path.basename(d['image_meta_dict']['filename_or_obj']).replace(\".nii\", \"\")\n            if uid not in empty_label_uids:\n                empty_label_uids.append(uid)\n        return d\n\nclass ApplyCLAHE(MapTransform):\n    def __init__(self, keys, clip_limit=3.5, tile_grid_size=(8, 8)):\n        super().__init__(keys)\n        self.clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n    def __call__(self, data):\n        d = dict(data)\n        for key in self.keys:\n            img = d[key]; img_uint8 = (img * 255).astype(np.uint8); clahe_slices = []\n            for i in range(img_uint8.shape[-1]):\n                clahe_slice = self.clahe.apply(img_uint8[0, :, :, i])\n                clahe_slices.append(clahe_slice)\n            clahe_img = np.stack(clahe_slices, axis=-1)\n            d[key] = (clahe_img / 255.0)[np.newaxis, ...].astype(np.float32)\n        return d\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Phase 4: Defining Transforms and DataLoaders\n# ==============================================================================\nprint(\"\\n--- Phase 4: Defining Data Transforms and DataLoaders ---\")\ntrain_transforms = Compose([\n    LoadImaged(keys=[\"image\", \"label\"]), EnsureChannelFirstd(keys=[\"image\", \"label\"]),\n    Orientationd(keys=[\"image\", \"label\"], axcodes=\"RAS\"),\n    Spacingd(keys=[\"image\", \"label\"], pixdim=(0.75, 0.75, 0.75), mode=(\"bilinear\", \"nearest\")),\n    ModalityNormalize(keys=[\"image\"]),\n    ApplyCLAHE(keys=[\"image\"]),  # <-- ADDED\n    CropForegroundd(keys=[\"image\", \"label\"], source_key=\"image\"),\n    Lambdad(keys=\"label\", func=lambda x: (x > 0).astype(x.dtype)),\n    LogEmptyLabelsd(keys=[\"image\", \"label\"]),\n    SpatialPadd(keys=[\"image\", \"label\"], spatial_size=PATCH_SIZE, method=\"end\"),\n    RandCropByPosNegLabeld(\n        keys=[\"image\", \"label\"], label_key=\"label\", spatial_size=PATCH_SIZE,\n        pos=1, neg=1, num_samples=4, image_key=\"image\", image_threshold=0,\n    ),\n    ToTensord(keys=[\"image\", \"label\"]), SelectItemsd(keys=[\"image\", \"label\"]),\n])\nval_transforms = Compose([\n    LoadImaged(keys=[\"image\", \"label\"]), EnsureChannelFirstd(keys=[\"image\", \"label\"]),\n    Orientationd(keys=[\"image\", \"label\"], axcodes=\"RAS\"),\n    Spacingd(keys=[\"image\", \"label\"], pixdim=(0.75, 0.75, 0.75), mode=(\"bilinear\", \"nearest\")),\n    ModalityNormalize(keys=[\"image\"]),\n    ApplyCLAHE(keys=[\"image\"]),  # <-- ADDED\n    CropForegroundd(keys=[\"image\", \"label\"], source_key=\"image\"),\n    Lambdad(keys=\"label\", func=lambda x: (x > 0).astype(x.dtype)),\n    DivisiblePadd(keys=[\"image\", \"label\"], k=32),\n    ToTensord(keys=[\"image\", \"label\"]), SelectItemsd(keys=[\"image\", \"label\"]),\n])\nprint(f\"Creating cache datasets (caching {CACHE_RATE*100}% of data to RAM)...\")\ntrain_ds = CacheDataset(data=train_files, transform=train_transforms, cache_rate=CACHE_RATE)\nval_ds = CacheDataset(data=val_files, transform=val_transforms, cache_rate=CACHE_RATE)\ntrain_loader = DataLoader(train_ds, batch_size=4, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_ds, batch_size=1, shuffle=False, num_workers=2)\nprint(\"DataLoaders are ready.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Phase 5: Model Definition, Loading, and Adaptation (Corrected)\n# ==============================================================================\nprint(\"\\n--- Phase 5: Defining and Preparing SwinUNETR Model ---\")\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# --- THIS IS THE CORRECTED PART ---\n# Use the 'img_size' argument required by the older MONAI version\nmodel = SwinUNETR(\n    img_size=PATCH_SIZE,  # Use the PATCH_SIZE variable (64, 64, 64)\n    in_channels=1,\n    out_channels=3,\n    feature_size=48,\n    use_checkpoint=False\n)\n# --- END OF CORRECTION ---\n\nweights_url = \"https://github.com/Project-MONAI/MONAI-extra-test-data/releases/download/0.8.1/model_swinvit.pt\"\ncheckpoint = load_state_dict_from_url(weights_url, map_location='cpu')\nmodel.load_state_dict(checkpoint['state_dict'], strict=False)\n\nin_features = model.out.conv.conv.in_channels\nmodel.out.conv = nn.Conv3d(in_features, out_channels=2, kernel_size=1)\n\nmodel.to(device)\nprint(\"Model adapted and moved to GPU.\")\n\nfor param in model.parameters():\n    param.requires_grad = False\n    \nlayers_to_unfreeze = ['decoder1', 'out']\nfor name, param in model.named_parameters():\n    for layer_name in layers_to_unfreeze:\n        if layer_name in name:\n            param.requires_grad = True\n            break\n            \ntrainable_params = [p for p in model.parameters() if p.requires_grad]\nprint(\"Model fine-tuning layers are set.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Phase 6: Loss, Optimizer, and Training Loop (with Validation Progress Bar)\n# ==============================================================================\nprint(\"\\n--- Phase 6: Starting Training Loop ---\")\n\nimport gc\nfrom tqdm import tqdm\nfrom monai.metrics import DiceMetric\nfrom monai.data import decollate_batch\nfrom monai.transforms import AsDiscrete\n\n# (The SoftDiceFocalLoss class definition remains the same)\nclass SoftDiceFocalLoss(nn.Module):\n    def __init__(self, weight_dice=0.6, weight_focal=0.4):\n        super().__init__()\n        self.dice_loss = DiceLoss(to_onehot_y=True, softmax=True)\n        self.focal_loss = FocalLoss(to_onehot_y=True, gamma=2.0)\n        self.weight_dice = weight_dice\n        self.weight_focal = weight_focal\n    \n    def forward(self, y_pred, y_true):\n        loss_d = self.dice_loss(y_pred, y_true)\n        loss_f = self.focal_loss(y_pred, y_true)\n        total_loss = (self.weight_dice * loss_d) + (self.weight_focal * loss_f)\n        return total_loss, loss_d.item(), loss_f.item()\n\nloss_function = SoftDiceFocalLoss()\noptimizer = torch.optim.AdamW(trainable_params, lr=1e-4, weight_decay=1e-5)\ndice_metric = DiceMetric(include_background=False, reduction=\"mean\")\npost_pred = AsDiscrete(argmax=True, to_onehot=2)\npost_label = AsDiscrete(to_onehot=2)\nbest_metric = -1\nbest_metric_epoch = -1\nepoch_loss_values = []\nmetric_values = []\n\nfor epoch in range(NUM_EPOCHS):\n    print(\"-\" * 20, f\"\\nEpoch {epoch + 1}/{NUM_EPOCHS}\")\n    model.train()\n    epoch_loss = 0\n    \n    progress_bar = tqdm(train_loader, desc=\"Training\", colour=\"green\")\n    for step, batch_data in enumerate(progress_bar, 1):\n        inputs, labels = batch_data[\"image\"].to(device), batch_data[\"label\"].to(device)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        total_loss, dice_val, focal_val = loss_function(outputs, labels)\n        total_loss.backward()\n        optimizer.step()\n        epoch_loss += total_loss.item()\n        progress_bar.set_postfix({\"loss\": f\"{total_loss.item():.4f}\", \"dice\": f\"{dice_val:.4f}\", \"focal\": f\"{focal_val:.4f}\"})\n        \n    epoch_loss /= step\n    print(f\"Average Training Loss: {epoch_loss:.4f}\")\n    epoch_loss_values.append(epoch_loss)\n    \n    model.eval()\n    with torch.no_grad():\n        # --- THIS IS THE MODIFIED PART ---\n        # Wrap the validation data loader with tqdm for a progress bar\n        val_progress_bar = tqdm(val_loader, desc=\"Validation\", colour=\"blue\")\n        for val_data in val_progress_bar:\n            val_inputs, val_labels = val_data[\"image\"].to(device), val_data[\"label\"].to(device)\n            val_outputs = sliding_window_inference(inputs=val_inputs, roi_size=PATCH_SIZE, sw_batch_size=4, predictor=model)\n            val_outputs = [post_pred(i) for i in decollate_batch(val_outputs)]\n            val_labels = [post_label(i) for i in decollate_batch(val_labels)]\n            dice_metric(y_pred=val_outputs, y=val_labels)\n            \n        metric = dice_metric.aggregate().item()\n        dice_metric.reset()\n        \n    print(f\"Mean Validation Dice Score: {metric:.4f}\")\n    metric_values.append(metric)\n    \n    if metric > best_metric:\n        best_metric = metric\n        best_metric_epoch = epoch + 1\n        torch.save(model.state_dict(), BEST_MODEL_FILENAME)\n        print(f\"  >> New best model saved as {BEST_MODEL_FILENAME}!\")\n        \nprint(f\"\\n--- Training Finished ---\\nBest Dice Score: {best_metric:.4f} at Epoch {best_metric_epoch}\")\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Phase 7: Plotting Metrics\n# ==============================================================================\nprint(\"\\n--- Phase 7: Plotting Metrics ---\")\nplt.figure(\"train\", (12, 6))\nplt.subplot(1, 2, 1)\nplt.title(\"Epoch Average Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nx = [i + 1 for i in range(len(epoch_loss_values))]\nplt.plot(x, epoch_loss_values)\nplt.grid(True)\n\nplt.subplot(1, 2, 2)\nplt.title(\"Validation Mean Dice Score\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Dice Score\")\nx = [i + 1 for i in range(len(metric_values))]\nplt.plot(x, metric_values)\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Phase 8: Log Empty Label Information\n# ==============================================================================\nprint(\"\\n--- Phase 8: Logging Empty Label Information ---\")\nunique_empty_count = len(set(empty_label_uids))\nprint(f\"Found {unique_empty_count} unique training samples with empty labels.\")\noutput_log_file = \"empty_label_patients.txt\"\nwith open(output_log_file, 'w') as f:\n    for uid in sorted(list(set(empty_label_uids))):\n        f.write(f\"{uid}\\n\")\nprint(f\"A list of these patient UIDs has been saved to: {output_log_file}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==============================================================================\n# Phase 9: Inference and Saving Prediction\n# ==============================================================================\nprint(\"\\n--- Phase 9: Running Inference on a Sample ---\")\nif best_metric > -1:\n    model.load_state_dict(torch.load(BEST_MODEL_FILENAME))\n    model.eval()\n    sample_data = val_files[0]\n    original_image_path = sample_data[\"image\"]\n    with torch.no_grad():\n        input_data = val_transforms(sample_data)\n        input_tensor = input_data[\"image\"].unsqueeze(0).to(device)\n        prediction_logits = sliding_window_inference(inputs=input_tensor, roi_size=PATCH_SIZE, sw_batch_size=4, predictor=model)\n        predicted_mask_tensor = torch.argmax(prediction_logits, dim=1).squeeze(0)\n    \n    prediction_numpy = predicted_mask_tensor.cpu().numpy().astype(np.uint8)\n    original_image_nii = nib.load(original_image_path)\n    prediction_nii = nib.Nifti1Image(prediction_numpy, original_image_nii.affine)\n    output_filename = \"predicted_segmentation.nii\"\n    nib.save(prediction_nii, output_filename)\n    print(f\"Prediction saved to: {output_filename}\")\nelse:\n    print(\"Skipping inference because no model was saved.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}