{"metadata":{"kernelspec":{"display_name":"py11","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.11.14"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom PIL import Image, ImageSequence\nimport tifffile as tiff\nimport zipfile\nfrom scipy.ndimage import gaussian_filter","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2025-12-15T18:20:24.026219Z","iopub.status.busy":"2025-12-15T18:20:24.025522Z","iopub.status.idle":"2025-12-15T18:20:24.031102Z","shell.execute_reply":"2025-12-15T18:20:24.029969Z","shell.execute_reply.started":"2025-12-15T18:20:24.026191Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Auto-detect environment (Kaggle vs Local)\nimport os\n\n# Check if running on Kaggle\nif os.path.exists(\"/kaggle/input\"):\n    DATA_ROOT = \"/kaggle/input/vesuvius-challenge-surface-detection\"\n    print(\" Running on Kaggle\")\nelse:\n    # Local execution\n    DATA_ROOT = \"vesuvius-challenge-surface-detection\"\n    print(\" Running locally\")\n\n# Verify the path exists\nif not os.path.exists(DATA_ROOT):\n    print(f\"⚠️  WARNING: DATA_ROOT '{DATA_ROOT}' not found!\")\n    print(f\"Current directory: {os.getcwd()}\")\n    print(f\"\\nAvailable paths:\")\n    if os.path.exists(\"/kaggle/input\"):\n        print(\"Kaggle inputs:\", os.listdir(\"/kaggle/input\"))\n    else:\n        print(\"Local directory:\", os.listdir(\".\"))\nelse:\n    print(f\"✓ DATA_ROOT found: {DATA_ROOT}\")\n    print(f\"  Contents: {os.listdir(DATA_ROOT)}\")\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"\\n  Device: {DEVICE}\")\n\n# Training parameters\nBATCH_SIZE = 2          # Increased batch size for faster training\nEPOCHS = 5              # Reduced from 15 - should be enough to see results\nLR = 1e-4\n\n# 3D Volume parameters (optimized for speed)\nPATCH_SIZE = (64, 64, 64)    # Reduced from (96,96,96) - much faster!\nROI_SIZE = (96, 96, 96)      # Reduced from (128,128,128)\nNUM_CLASSES = 3              # Background, recto, verso\nOVERLAP = 0.25               # Reduced from 0.5 - faster inference!\n\n# Prediction threshold\nTHRESHOLD = 0.5\n\n# Checkpointing\nCHECKPOINT_DIR = \"/kaggle/working\" if os.path.exists(\"/kaggle\") else \"checkpoints\"\nos.makedirs(CHECKPOINT_DIR, exist_ok=True)","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.032566Z","iopub.status.busy":"2025-12-15T18:20:24.032297Z","iopub.status.idle":"2025-12-15T18:20:24.047479Z","shell.execute_reply":"2025-12-15T18:20:24.046522Z","shell.execute_reply.started":"2025-12-15T18:20:24.032538Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check GPU availability\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"CUDA version: {torch.version.cuda}\")\n    print(f\"GPU count: {torch.cuda.device_count()}\")\n    print(f\"GPU name: {torch.cuda.get_device_name(0)}\")\nelse:\n    print(\"\\nNo GPU detected. Options:\")\n    print(\"1. Install PyTorch with CUDA: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121\")\n    print(\"2. Or upload this notebook to Kaggle to use free GPU\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_volume(path):\n    \"\"\"\n    Read LZW-compressed multi-page TIFF using PIL\n    Returns (D, H, W) float32 with intensity normalization [0, 1]\n    \"\"\"\n    img = Image.open(path)\n\n    slices = []\n    for page in ImageSequence.Iterator(img):\n        slices.append(np.array(page, dtype=np.float32))\n\n    vol = np.stack(slices, axis=0)  # (D, H, W)\n    \n    # Intensity normalization to [0, 1] range (similar to ScaleIntensityRange in reference)\n    vol = np.clip(vol, 0, 255) / 255.0\n    \n    return vol\n\ndef normalize_volume(vol):\n    \"\"\"\n    Z-score normalization for training\n    \"\"\"\n    mean = vol.mean()\n    std = vol.std()\n    if std > 1e-6:\n        vol = (vol - mean) / std\n    return vol","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.049250Z","iopub.status.busy":"2025-12-15T18:20:24.048760Z","iopub.status.idle":"2025-12-15T18:20:24.064366Z","shell.execute_reply":"2025-12-15T18:20:24.063556Z","shell.execute_reply.started":"2025-12-15T18:20:24.049221Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VolumeDataset3D(Dataset):\n    \"\"\"\n    3D Volume Dataset that extracts random 3D patches during training\n    \"\"\"\n    def __init__(self, image_ids, patch_size=(96, 96, 96), is_train=True):\n        self.image_ids = image_ids\n        self.patch_size = patch_size\n        self.is_train = is_train\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        img_id = self.image_ids[idx]\n\n        # Load full volumes\n        img = load_volume(f\"{DATA_ROOT}/train_images/{img_id}.tif\")\n        lbl = load_volume(f\"{DATA_ROOT}/train_labels/{img_id}.tif\")\n        \n        # Apply z-score normalization for training\n        img = normalize_volume(img)\n\n        D, H, W = img.shape\n        pd, ph, pw = self.patch_size\n\n        # Random crop for training, center crop for validation\n        if self.is_train and D > pd and H > ph and W > pw:\n            # Random 3D patch extraction\n            d_start = np.random.randint(0, D - pd + 1)\n            h_start = np.random.randint(0, H - ph + 1)\n            w_start = np.random.randint(0, W - pw + 1)\n        else:\n            # Center crop\n            d_start = max(0, (D - pd) // 2)\n            h_start = max(0, (H - ph) // 2)\n            w_start = max(0, (W - pw) // 2)\n\n        # Extract patches\n        d_end = min(d_start + pd, D)\n        h_end = min(h_start + ph, H)\n        w_end = min(w_start + pw, W)\n\n        img_patch = img[d_start:d_end, h_start:h_end, w_start:w_end]\n        lbl_patch = lbl[d_start:d_end, h_start:h_end, w_start:w_end]\n\n        # Pad if necessary\n        if img_patch.shape != self.patch_size:\n            pad_d = pd - img_patch.shape[0]\n            pad_h = ph - img_patch.shape[1]\n            pad_w = pw - img_patch.shape[2]\n            \n            img_patch = np.pad(img_patch, (\n                (0, pad_d), (0, pad_h), (0, pad_w)\n            ), mode='constant', constant_values=0)\n            \n            lbl_patch = np.pad(lbl_patch, (\n                (0, pad_d), (0, pad_h), (0, pad_w)\n            ), mode='constant', constant_values=0)\n\n        # Add channel dimension: (1, D, H, W)\n        x = img_patch[None, ...]\n        \n        # Convert labels to class indices (0: background, 1+: foreground classes)\n        # Assuming binary labels, convert to 2-class (or 3-class if needed)\n        y = (lbl_patch > 0.5).astype(np.int64)  # Binary: 0 or 1\n\n        return torch.tensor(x, dtype=torch.float32), torch.tensor(y, dtype=torch.long)","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.085934Z","iopub.status.busy":"2025-12-15T18:20:24.085677Z","iopub.status.idle":"2025-12-15T18:20:24.103465Z","shell.execute_reply":"2025-12-15T18:20:24.102384Z","shell.execute_reply.started":"2025-12-15T18:20:24.085914Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet3D(nn.Module):\n    \"\"\"\n    Improved 3D U-Net for volumetric segmentation\n    Based on the architecture style from the reference\n    \"\"\"\n    def __init__(self, in_channels=1, num_classes=2, base_filters=32):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._conv_block(in_channels, base_filters)\n        self.pool1 = nn.MaxPool3d(2)\n        \n        self.enc2 = self._conv_block(base_filters, base_filters * 2)\n        self.pool2 = nn.MaxPool3d(2)\n        \n        self.enc3 = self._conv_block(base_filters * 2, base_filters * 4)\n        self.pool3 = nn.MaxPool3d(2)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(base_filters * 4, base_filters * 8)\n        \n        # Decoder\n        self.upconv3 = nn.ConvTranspose3d(base_filters * 8, base_filters * 4, 2, stride=2)\n        self.dec3 = self._conv_block(base_filters * 8, base_filters * 4)\n        \n        self.upconv2 = nn.ConvTranspose3d(base_filters * 4, base_filters * 2, 2, stride=2)\n        self.dec2 = self._conv_block(base_filters * 4, base_filters * 2)\n        \n        self.upconv1 = nn.ConvTranspose3d(base_filters * 2, base_filters, 2, stride=2)\n        self.dec1 = self._conv_block(base_filters * 2, base_filters)\n        \n        # Output\n        self.out = nn.Conv3d(base_filters, num_classes, 1)\n    \n    def _conv_block(self, in_ch, out_ch):\n        return nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        \n        # Bottleneck\n        b = self.bottleneck(self.pool3(e3))\n        \n        # Decoder with skip connections\n        d3 = self.upconv3(b)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.upconv2(d3)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.upconv1(d2)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        out = self.out(d1)\n        return out","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.085934Z","iopub.status.busy":"2025-12-15T18:20:24.085677Z","iopub.status.idle":"2025-12-15T18:20:24.103465Z","shell.execute_reply":"2025-12-15T18:20:24.102384Z","shell.execute_reply.started":"2025-12-15T18:20:24.085914Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_loss(pred, target, smooth=1e-6):\n    \"\"\"\n    Dice loss for multi-class segmentation\n    pred: (B, C, D, H, W) logits\n    target: (B, D, H, W) class indices\n    \"\"\"\n    # Convert logits to probabilities\n    pred = F.softmax(pred, dim=1)\n    \n    # One-hot encode targets\n    num_classes = pred.shape[1]\n    target_one_hot = F.one_hot(target, num_classes=num_classes)\n    target_one_hot = target_one_hot.permute(0, 4, 1, 2, 3).float()\n    \n    # Compute Dice for each class\n    dice_scores = []\n    for c in range(num_classes):\n        pred_c = pred[:, c]\n        target_c = target_one_hot[:, c]\n        \n        intersection = (pred_c * target_c).sum()\n        union = pred_c.sum() + target_c.sum()\n        \n        dice = (2.0 * intersection + smooth) / (union + smooth)\n        dice_scores.append(dice)\n    \n    # Average across classes (excluding background if needed)\n    dice_loss = 1.0 - torch.stack(dice_scores).mean()\n    return dice_loss\n\ndef combined_loss(pred, target):\n    \"\"\"\n    Combined Cross Entropy + Dice Loss\n    \"\"\"\n    ce_loss = F.cross_entropy(pred, target)\n    d_loss = dice_loss(pred, target)\n    return 0.5 * ce_loss + 0.5 * d_loss","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.104580Z","iopub.status.busy":"2025-12-15T18:20:24.104341Z","iopub.status.idle":"2025-12-15T18:20:24.119665Z","shell.execute_reply":"2025-12-15T18:20:24.118554Z","shell.execute_reply.started":"2025-12-15T18:20:24.104554Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_ids = sorted([\n    f.replace(\".tif\", \"\")\n    for f in os.listdir(f\"{DATA_ROOT}/train_images\")\n])\n\nprint(f\"Training on {len(image_ids)} images\")\n\ndataset = VolumeDataset3D(image_ids, patch_size=PATCH_SIZE, is_train=True)\n\nloader = DataLoader(\n    dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=0,\n    pin_memory=True if DEVICE == \"cuda\" else False,\n)","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.120868Z","iopub.status.busy":"2025-12-15T18:20:24.120594Z","iopub.status.idle":"2025-12-15T18:20:24.142070Z","shell.execute_reply":"2025-12-15T18:20:24.141151Z","shell.execute_reply.started":"2025-12-15T18:20:24.120839Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet3D(in_channels=1, num_classes=NUM_CLASSES, base_filters=16).to(DEVICE)  # Reduced from 32\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\nprint(f\"Model parameters: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M\")\nprint(f\" Training optimizations: Smaller model + patches = ~5-10x faster!\")\n\nbest_loss = float('inf')\npatience_counter = 0\nPATIENCE = 3  # Early stopping patience\n\nfor epoch in range(EPOCHS):\n    model.train()\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    for x, y in tqdm(loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n\n        optimizer.zero_grad()\n        pred = model(x)\n        \n        loss = combined_loss(pred, y)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n        \n        # Calculate accuracy\n        pred_classes = pred.argmax(dim=1)\n        correct += (pred_classes == y).sum().item()\n        total += y.numel()\n\n    avg_loss = total_loss / len(loader)\n    accuracy = 100.0 * correct / total\n    print(f\"Epoch {epoch+1} - Loss: {avg_loss:.4f}, Accuracy: {accuracy:.2f}%\")\n    \n    # Save checkpoint every epoch\n    checkpoint_path = f\"{CHECKPOINT_DIR}/model_epoch_{epoch+1}.pth\"\n    torch.save({\n        'epoch': epoch + 1,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': avg_loss,\n    }, checkpoint_path)\n    print(f\"  💾 Checkpoint saved: {checkpoint_path}\")\n    \n    # Early stopping check\n    if avg_loss < best_loss:\n        best_loss = avg_loss\n        patience_counter = 0\n        # Save best model\n        torch.save(model.state_dict(), f\"{CHECKPOINT_DIR}/best_model.pth\")\n        print(f\"  ⭐ New best model saved!\")\n    else:\n        patience_counter += 1\n        if patience_counter >= PATIENCE:\n            print(f\"   Early stopping triggered (no improvement for {PATIENCE} epochs)\")\n            break\n\nprint(f\"\\n✅ Training complete! Best loss: {best_loss:.4f}\")","metadata":{"execution":{"iopub.execute_input":"2025-12-15T18:20:24.143315Z","iopub.status.busy":"2025-12-15T18:20:24.143012Z","iopub.status.idle":"2025-12-15T18:20:43.493286Z","shell.execute_reply":"2025-12-15T18:20:43.491722Z","shell.execute_reply.started":"2025-12-15T18:20:24.143286Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sliding_window_inference(model, volume, roi_size=(128, 128, 128), overlap=0.5, mode='gaussian'):\n    \"\"\"\n    Sliding window inference for 3D volumes (similar to reference implementation)\n    \n    Args:\n        model: trained model\n        volume: input volume (D, H, W)\n        roi_size: size of sliding window\n        overlap: overlap ratio between windows\n        mode: 'gaussian' for gaussian weighting, 'constant' for uniform\n    \n    Returns:\n        prediction: (D, H, W) class indices\n    \"\"\"\n    model.eval()\n    \n    D, H, W = volume.shape\n    rd, rh, rw = roi_size\n    \n    # Calculate stride based on overlap\n    stride_d = int(rd * (1 - overlap))\n    stride_h = int(rh * (1 - overlap))\n    stride_w = int(rw * (1 - overlap))\n    \n    # Initialize output and weight maps\n    output = np.zeros((NUM_CLASSES, D, H, W), dtype=np.float32)\n    weights = np.zeros((D, H, W), dtype=np.float32)\n    \n    # Use constant weighting for speed (gaussian is slow!)\n    importance_map = np.ones(roi_size, dtype=np.float32)\n    \n    # Sliding window loop\n    with torch.no_grad():\n        for d_start in range(0, D, stride_d):\n            for h_start in range(0, H, stride_h):\n                for w_start in range(0, W, stride_w):\n                    # Calculate patch boundaries\n                    d_end = min(d_start + rd, D)\n                    h_end = min(h_start + rh, H)\n                    w_end = min(w_start + rw, W)\n                    \n                    # Adjust start if we're at the edge\n                    if d_end == D:\n                        d_start = max(0, D - rd)\n                        d_end = D\n                    if h_end == H:\n                        h_start = max(0, H - rh)\n                        h_end = H\n                    if w_end == W:\n                        w_start = max(0, W - rw)\n                        w_end = W\n                    \n                    # Extract patch\n                    patch = volume[d_start:d_end, h_start:h_end, w_start:w_end]\n                    \n                    # Pad if necessary\n                    actual_shape = patch.shape\n                    if patch.shape != roi_size:\n                        pad_d = rd - patch.shape[0]\n                        pad_h = rh - patch.shape[1]\n                        pad_w = rw - patch.shape[2]\n                        patch = np.pad(patch, ((0, pad_d), (0, pad_h), (0, pad_w)), \n                                     mode='constant', constant_values=0)\n                    \n                    # Predict\n                    x = torch.tensor(patch[None, None, ...], dtype=torch.float32).to(DEVICE)\n                    pred = model(x)\n                    pred = F.softmax(pred, dim=1)[0].cpu().numpy()\n                    \n                    # Remove padding\n                    pred = pred[:, :actual_shape[0], :actual_shape[1], :actual_shape[2]]\n                    \n                    # Add to output with importance weighting\n                    imp_crop = importance_map[:actual_shape[0], :actual_shape[1], :actual_shape[2]]\n                    output[:, d_start:d_end, h_start:h_end, w_start:w_end] += pred * imp_crop\n                    weights[d_start:d_end, h_start:h_end, w_start:w_end] += imp_crop\n    \n    # Normalize by weights\n    weights[weights == 0] = 1\n    output = output / weights[None, ...]\n    \n    # Get class predictions\n    prediction = output.argmax(axis=0).astype(np.uint8)\n    \n    return prediction","metadata":{"execution":{"iopub.status.busy":"2025-12-15T18:20:43.493836Z","iopub.status.idle":"2025-12-15T18:20:43.494092Z","shell.execute_reply":"2025-12-15T18:20:43.493986Z","shell.execute_reply.started":"2025-12-15T18:20:43.493975Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Use /kaggle/working on Kaggle, local 'submission' folder otherwise\noutput_dir = \"/kaggle/working\" if os.path.exists(\"/kaggle\") else \"submission\"\nos.makedirs(output_dir, exist_ok=True)\n\ntest_ids = [\n    f.replace(\".tif\", \"\")\n    for f in os.listdir(f\"{DATA_ROOT}/test_images\")\n]\n\nprint(f\"Predicting {len(test_ids)} test volumes with sliding window inference...\")\n\nfor img_id in tqdm(test_ids, desc=\"Predicting test volumes\"):\n    # Load and normalize volume\n    vol = load_volume(f\"{DATA_ROOT}/test_images/{img_id}.tif\")\n    vol = normalize_volume(vol)\n    \n    # Sliding window inference (using constant mode for speed)\n    pred = sliding_window_inference(\n        model, \n        vol, \n        roi_size=ROI_SIZE, \n        overlap=OVERLAP,\n        mode='constant'  # Much faster than gaussian!\n    )\n    \n    # Save prediction\n    tiff.imwrite(f\"{output_dir}/{img_id}.tif\", pred.astype(np.uint8))\n\nprint(f\" Predictions saved to {output_dir}\")","metadata":{"execution":{"iopub.status.busy":"2025-12-15T18:20:43.495638Z","iopub.status.idle":"2025-12-15T18:20:43.495975Z","shell.execute_reply":"2025-12-15T18:20:43.495837Z","shell.execute_reply.started":"2025-12-15T18:20:43.495825Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create submission zip\noutput_dir = \"/kaggle/working\" if os.path.exists(\"/kaggle\") else \"submission\"\nzip_path = \"submission.zip\" if not os.path.exists(\"/kaggle\") else \"/kaggle/working/submission.zip\"\n\nwith zipfile.ZipFile(zip_path, \"w\") as z:\n    for f in os.listdir(output_dir):\n        if f.endswith(\".tif\"):\n            z.write(f\"{output_dir}/{f}\", arcname=f)\n\nprint(f\"{zip_path} created\")\n","metadata":{"execution":{"iopub.status.busy":"2025-12-15T18:20:43.496899Z","iopub.status.idle":"2025-12-15T18:20:43.497188Z","shell.execute_reply":"2025-12-15T18:20:43.497053Z","shell.execute_reply.started":"2025-12-15T18:20:43.497038Z"},"trusted":true},"outputs":[],"execution_count":null}]}