{"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.8.10"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":14443416,"sourceType":"competition"},{"sourceId":278863510,"sourceType":"kernelVersion"},{"sourceId":279270726,"sourceType":"kernelVersion"},{"sourceId":280878254,"sourceType":"kernelVersion"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Vesuvius Challenge - 3D UNet Inference\n\nPerforms inference on test data using trained MONAI 3D UNet models.\n\n**Pipeline:**\n1. Load test volumes (.tif files, any size)\n2. Downsample to model size (256³) \n3. Run patch-based inference with ensemble\n4. Upsample predictions back to actual input size\n5. Save predictions as uint8 .tif files\n6. Zip all predictions for submission","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y tensorflow protobuf\n!pip install --no-deps /kaggle/input/wheels-vesuvius/monai-1.5.1-py3-none-any.whl\n!pip install --no-deps /kaggle/input/wheels-vesuvius/imagecodecs-2025.11.11-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nfrom pathlib import Path\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nfrom monai.networks.nets import UNet\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:08.486799Z","iopub.execute_input":"2025-11-22T19:05:08.487497Z","iopub.status.idle":"2025-11-22T19:05:08.491822Z","shell.execute_reply.started":"2025-11-22T19:05:08.487471Z","shell.execute_reply":"2025-11-22T19:05:08.490892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # Data directories\n    TEST_IMG_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection/test_images\")\n    MODEL_DIR = Path(\"/kaggle/input/vesuvius-eda-monai-3d-unet-baseline\")\n    \n    # Volume sizes\n    MODEL_SIZE = (256, 256, 256)     # Model input size\n    \n    # Inference settings\n    PATCH_SIZE = (128, 128, 128)\n    BATCH_SIZE = 8\n    FOLDS = 5\n    THRESHOLD = 0.6\n    \n    # Output settings\n    OUTPUT_DIR = Path(\"./predictions\")\n    SAVE_VISUALIZATIONS = True  # Set to True to save visualization images\n    \n    # Device\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Create output directories\nCFG.OUTPUT_DIR.mkdir(exist_ok=True, parents=True)\n(CFG.OUTPUT_DIR / \"submission_tifs\").mkdir(exist_ok=True, parents=True)\n\nprint(f\"Device: {CFG.DEVICE}\")\nprint(f\"Threshold: {CFG.THRESHOLD}\")\nprint(f\"Visualizations: {CFG.SAVE_VISUALIZATIONS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:06:27.799138Z","iopub.execute_input":"2025-11-22T19:06:27.799434Z","iopub.status.idle":"2025-11-22T19:06:27.806085Z","shell.execute_reply.started":"2025-11-22T19:06:27.799414Z","shell.execute_reply":"2025-11-22T19:06:27.805301Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utility Functions","metadata":{}},{"cell_type":"code","source":"def load_array(path, fmt):\n    \"\"\"Load array from various formats\"\"\"\n    if fmt == \"tiff\" or fmt == \"tif\":\n        return tiff.imread(path)\n    elif fmt == \"npy\":\n        return np.load(path)\n    elif fmt == \"npz\":\n        return np.load(path)[\"arr_0\"]\n    elif fmt == \"rle\":\n        rle = np.load(path)\n        shape = tuple(rle[\"shape\"])\n        vals = rle[\"vals\"]\n        runs = rle[\"runs\"]\n        flat = np.repeat(vals, runs)\n        return flat.reshape(shape)\n    else:\n        raise ValueError(f\"Unsupported format: {fmt}\")\n\n\ndef normalize_volume(volume):\n    \"\"\"Normalize volume to zero mean and unit std\"\"\"\n    volume = volume.astype(np.float32)\n    mean = volume.mean()\n    std = volume.std()\n    return (volume - mean) / (std + 1e-6)\n\n\ndef downsample_volume(volume, target_size):\n    \"\"\"Downsample volume from original size to target size using trilinear interpolation\"\"\"\n    import torch.nn.functional as F\n    \n    # Convert to tensor\n    vol_tensor = torch.from_numpy(volume).float().unsqueeze(0).unsqueeze(0)\n    \n    # Downsample using trilinear interpolation\n    downsampled = F.interpolate(\n        vol_tensor,\n        size=target_size,\n        mode='trilinear',\n        align_corners=False\n    )\n    \n    return downsampled.squeeze(0).squeeze(0).numpy()\n\n\ndef upsample_volume(volume, target_size):\n    \"\"\"Upsample volume from model size back to original size using trilinear interpolation\"\"\"\n    import torch.nn.functional as F\n    \n    # Convert to tensor\n    vol_tensor = torch.from_numpy(volume).float().unsqueeze(0).unsqueeze(0)\n    \n    # Upsample using trilinear interpolation\n    upsampled = F.interpolate(\n        vol_tensor,\n        size=target_size,\n        mode='trilinear',\n        align_corners=False\n    )\n    \n    return upsampled.squeeze(0).squeeze(0).numpy()\n\n\ndef extract_all_patches(volume, patch_size):\n    \"\"\"Extract all non-overlapping patches from a volume\"\"\"\n    D, H, W = volume.shape\n    pd, ph, pw = patch_size\n    patches = []\n    coords = []\n    \n    for z in range(0, D - pd + 1, pd):\n        for y in range(0, H - ph + 1, ph):\n            for x in range(0, W - pw + 1, pw):\n                patch = volume[z:z+pd, y:y+ph, x:x+pw]\n                patches.append(patch)\n                coords.append((z, y, x))\n    \n    return patches, coords\n\n\ndef stitch_patches(patches, coords, full_shape, patch_size):\n    \"\"\"Stitch patches back into a full volume with averaging for overlaps\"\"\"\n    recon = np.zeros(full_shape, dtype=np.float32)\n    counts = np.zeros(full_shape, dtype=np.float32)\n    pd, ph, pw = patch_size\n\n    for patch, (z, y, x) in zip(patches, coords):\n        recon[z:z+pd, y:y+ph, x:x+pw] += patch\n        counts[z:z+pd, y:y+ph, x:x+pw] += 1\n\n    recon /= np.maximum(counts, 1)\n    return recon\n\n\ndef rle_encode(mask):\n    \"\"\"Run-length encoding for submission\"\"\"\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:09.365666Z","iopub.execute_input":"2025-11-22T19:05:09.366005Z","iopub.status.idle":"2025-11-22T19:05:09.379169Z","shell.execute_reply.started":"2025-11-22T19:05:09.365957Z","shell.execute_reply":"2025-11-22T19:05:09.378243Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architecture","metadata":{}},{"cell_type":"code","source":"class Ves3DUNet(torch.nn.Module):\n    \"\"\"3D UNet model for surface detection\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.model = UNet(\n            spatial_dims=3,\n            in_channels=1,\n            out_channels=1,\n            channels=(16, 32, 64, 128),\n            strides=(2, 2, 2, 2),\n            num_res_units=3,\n            norm='batch'\n        )\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:09.778072Z","iopub.execute_input":"2025-11-22T19:05:09.778806Z","iopub.status.idle":"2025-11-22T19:05:09.7834Z","shell.execute_reply.started":"2025-11-22T19:05:09.778784Z","shell.execute_reply":"2025-11-22T19:05:09.782484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference Dataset","metadata":{}},{"cell_type":"code","source":"class InferenceDataset(Dataset):\n    \"\"\"Dataset for inference on test volumes\"\"\"\n    def __init__(self, image_paths, patch_size, model_size):\n        self.image_paths = image_paths\n        self.patch_size = patch_size\n        self.model_size = model_size\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        # Load original volume and get actual shape\n        vol = load_array(img_path, \"tif\")\n        actual_shape = vol.shape  # Get actual input shape\n        \n        # Downsample to model size\n        vol_downsampled = downsample_volume(vol, self.model_size)\n        vol_normalized = normalize_volume(vol_downsampled)\n        \n        # Extract patches\n        patches, coords = extract_all_patches(vol_normalized, self.patch_size)\n        patches_tensor = torch.stack([torch.from_numpy(p).unsqueeze(0) for p in patches])\n        \n        return {\n            'patches': patches_tensor,\n            'coords': coords,\n            'model_shape': vol_normalized.shape,\n            'actual_shape': actual_shape,\n            'filename': img_path.name\n        }\n\n\ndef custom_collate_fn(batch):\n    \"\"\"Custom collate function\"\"\"\n    item = batch[0]\n    return {\n        'patches': item['patches'].unsqueeze(0),\n        'coords': item['coords'],\n        'model_shape': torch.tensor([item['model_shape']]),\n        'actual_shape': torch.tensor([item['actual_shape']]),\n        'filename': [item['filename']]\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:14.85478Z","iopub.execute_input":"2025-11-22T19:05:14.855466Z","iopub.status.idle":"2025-11-22T19:05:14.864436Z","shell.execute_reply.started":"2025-11-22T19:05:14.855433Z","shell.execute_reply":"2025-11-22T19:05:14.863635Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Models","metadata":{}},{"cell_type":"code","source":"def load_models(model_dir, num_folds, device):\n    \"\"\"Load all fold models\"\"\"\n    models = []\n    \n    for fold in range(num_folds):\n        model_path = model_dir / f\"unet3d_fold{fold}.pth\"\n        \n        if not model_path.exists():\n            print(f\"Warning: Model for fold {fold} not found at {model_path}\")\n            continue\n            \n        model = Ves3DUNet().to(device)\n        model.load_state_dict(torch.load(model_path, map_location=device))\n        model.eval()\n        models.append(model)\n        print(f\"Loaded model for fold {fold}\")\n    \n    print(f\"\\nTotal models loaded: {len(models)}\")\n    return models\n\n\n# Load all trained models\nmodels = load_models(CFG.MODEL_DIR, CFG.FOLDS, CFG.DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:16.544907Z","iopub.execute_input":"2025-11-22T19:05:16.545571Z","iopub.status.idle":"2025-11-22T19:05:16.807184Z","shell.execute_reply.started":"2025-11-22T19:05:16.545546Z","shell.execute_reply":"2025-11-22T19:05:16.806337Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference Function","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef predict_volume(models, patches, coords, full_shape, patch_size, device):\n    \"\"\"Predict using ensemble of models\"\"\"\n    all_predictions = []\n    \n    # Get predictions from each model\n    for model in models:\n        model.eval()\n        \n        # Process patches in batches\n        batch_predictions = []\n        for i in range(0, len(patches), CFG.BATCH_SIZE):\n            batch = patches[i:i+CFG.BATCH_SIZE].to(device)\n            \n            with torch.amp.autocast('cuda'):\n                preds = model(batch)\n                preds = torch.sigmoid(preds).cpu().numpy()\n            \n            batch_predictions.append(preds)\n        \n        # Concatenate all batch predictions\n        all_batch_preds = np.concatenate(batch_predictions, axis=0)\n        \n        # Extract individual patches and stitch\n        pred_patches = [all_batch_preds[i, 0] for i in range(all_batch_preds.shape[0])]\n        pred_full = stitch_patches(pred_patches, coords, full_shape, patch_size)\n        \n        all_predictions.append(pred_full)\n    \n    # Ensemble: average predictions from all models\n    ensemble_pred = np.mean(all_predictions, axis=0)\n    \n    return ensemble_pred, all_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:18.139459Z","iopub.execute_input":"2025-11-22T19:05:18.139779Z","iopub.status.idle":"2025-11-22T19:05:18.146154Z","shell.execute_reply.started":"2025-11-22T19:05:18.139756Z","shell.execute_reply":"2025-11-22T19:05:18.145366Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Run Inference on Test Data","metadata":{}},{"cell_type":"code","source":"# Get test image files (.tif files)\nif CFG.TEST_IMG_DIR.exists():\n    test_files = sorted([f for f in CFG.TEST_IMG_DIR.glob(\"*.tif\")])\n    print(f\"Found {len(test_files)} test .tif files\")\n    \n    if len(test_files) > 0:\n        print(\"\\nTest files:\")\n        for f in test_files[:5]:\n            print(f\"  - {f.name}\")\n        if len(test_files) > 5:\n            print(f\"  ... and {len(test_files) - 5} more\")\nelse:\n    print(f\"Warning: Test directory not found at {CFG.TEST_IMG_DIR}\")\n    print(\"Please update CFG.TEST_IMG_DIR to point to the correct directory\")\n    test_files = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:20.292328Z","iopub.execute_input":"2025-11-22T19:05:20.292571Z","iopub.status.idle":"2025-11-22T19:05:20.298703Z","shell.execute_reply.started":"2025-11-22T19:05:20.292554Z","shell.execute_reply":"2025-11-22T19:05:20.298005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run inference and save predictions directly to .tif files\nif len(test_files) > 0 and len(models) > 0:\n    tif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"\n    processed_count = 0\n    \n    # Create dataset and dataloader\n    test_dataset = InferenceDataset(test_files, CFG.PATCH_SIZE, CFG.MODEL_SIZE)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=0, collate_fn=custom_collate_fn)\n    \n    print(\"\\nRunning inference...\\n\")\n    \n    for batch in tqdm(test_loader, desc=\"Processing volumes\"):\n        patches = batch['patches'][0]\n        coords = batch['coords']\n        model_shape = tuple(batch['model_shape'][0].numpy())\n        actual_shape = tuple(batch['actual_shape'][0].numpy())\n        filename = batch['filename'][0]\n        scroll_id = filename.replace('.tif', '')\n        \n        # Get ensemble prediction at model size (256^3)\n        ensemble_pred_256, _ = predict_volume(models, patches, coords, model_shape, CFG.PATCH_SIZE, CFG.DEVICE)\n        \n        # Upsample to actual input shape\n        ensemble_pred_actual = upsample_volume(ensemble_pred_256, actual_shape)\n        \n        # Create binary mask and ensure uint8 type\n        binary_mask = (ensemble_pred_actual > CFG.THRESHOLD).astype(np.uint8)\n        \n        # Save directly to .tif\n        tif_path = tif_dir / f\"{scroll_id}.tif\"\n        tiff.imwrite(tif_path, binary_mask)\n        processed_count += 1\n        \n        # Free memory\n        del ensemble_pred_256, ensemble_pred_actual, binary_mask\n    \n    print(f\"\\n✓ Inference complete! Saved {processed_count} .tif files to {tif_dir}\")\nelse:\n    print(\"No test files or models available.\")\n    processed_count = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:05:23.174525Z","iopub.execute_input":"2025-11-22T19:05:23.174853Z","iopub.status.idle":"2025-11-22T19:05:56.541758Z","shell.execute_reply.started":"2025-11-22T19:05:23.174832Z","shell.execute_reply":"2025-11-22T19:05:56.540997Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Submission","metadata":{}},{"cell_type":"code","source":"import zipfile\n\n# Zip all .tif files for submission\nif processed_count > 0:\n    tif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"\n    tif_files = sorted(tif_dir.glob(\"*.tif\"))\n    zip_path = \"submission.zip\"\n    \n    with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n        for tif_file in tif_files:\n            zipf.write(tif_file, tif_file.name)\n    \n    print(\"\\n=== Submission Created ===\")\n    print(f\"Files: {len(tif_files)} .tif files\")\n    print(f\"Zip: {zip_path}\")\n    print(f\"Threshold: {CFG.THRESHOLD}\")\nelse:\n    print(\"No predictions to zip.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:06:08.169365Z","iopub.execute_input":"2025-11-22T19:06:08.170103Z","iopub.status.idle":"2025-11-22T19:06:08.603549Z","shell.execute_reply.started":"2025-11-22T19:06:08.170078Z","shell.execute_reply":"2025-11-22T19:06:08.602861Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Preview Visualizations","metadata":{}},{"cell_type":"code","source":"# Display binary mask visualizations (if enabled)\nif CFG.SAVE_VISUALIZATIONS and processed_count > 0:\n    tif_dir = CFG.OUTPUT_DIR / \"submission_tifs\"\n    tif_files = sorted(tif_dir.glob(\"*.tif\"))[:3]  # Show first 3 files\n    \n    if tif_files:\n        print(\"\\n=== Binary Mask Previews ===\\n\")\n        \n        for tif_file in tif_files:\n            # Load the binary mask\n            binary_mask = tiff.imread(tif_file)\n            scroll_id = tif_file.stem\n            \n            # Show 2 slices\n            mid_slice = binary_mask.shape[0] // 2\n            fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n            \n            axes[0].imshow(binary_mask[mid_slice], cmap='gray')\n            axes[0].set_title(f'{scroll_id} - Slice {mid_slice}')\n            axes[0].axis('off')\n            \n            axes[1].imshow(binary_mask[mid_slice + binary_mask.shape[0]//4], cmap='gray')\n            axes[1].set_title(f'{scroll_id} - Slice {mid_slice + binary_mask.shape[0]//4}')\n            axes[1].axis('off')\n            \n            plt.tight_layout()\n            plt.show()\n            \n        print(f\"\\n✓ Displayed previews for {len(tif_files)} files\")\n    else:\n        print(\"No .tif files found for visualization.\")\nelse:\n    if not CFG.SAVE_VISUALIZATIONS:\n        print(\"Visualizations disabled (set CFG.SAVE_VISUALIZATIONS = True to enable)\")\n    else:\n        print(\"No predictions to visualize.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-22T19:06:31.289795Z","iopub.execute_input":"2025-11-22T19:06:31.290464Z","iopub.status.idle":"2025-11-22T19:06:31.599801Z","shell.execute_reply.started":"2025-11-22T19:06:31.290439Z","shell.execute_reply":"2025-11-22T19:06:31.598925Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Summary\n\n**Outputs:**\n- `predictions/submission_tifs/*.tif` - Binary uint8 masks (actual input size)\n- `predictions/submission.zip` - Zipped submission file\n\n**Next steps:**\nSubmit `submission.zip` to the competition","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}