{"cells": [{"cell_type": "markdown", "metadata": {}, "source": "# Vesuvius Challenge - V2+V3 Ensemble with TTA\n\n**Best local score:** 0.5550 (V2+V3 ensemble + TTA + t=0.25 + rm_small(200))\n\nConfiguration:\n- V2 weight: 1.0\n- V3 weight: 0.5\n- TTA: 8x flip augmentation\n- Threshold: 0.25\n- Post-processing: Remove small components (<200 voxels)"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Install imagecodecs from offline wheel (required for LZW-compressed TIFFs)\nimport subprocess\nimport sys\nfrom pathlib import Path\n\n# Check for offline wheel installer\nwheel_dir = Path(\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\")\nif wheel_dir.exists():\n    print(\"Installing imagecodecs from offline wheels...\")\n    subprocess.run([\n        sys.executable, \"-m\", \"pip\", \"install\", \"--quiet\",\n        \"--no-index\", \"--find-links\", str(wheel_dir),\n        \"imagecodecs\", \"tifffile\"\n    ], check=True)\n    print(\"Done!\")\n\nimport imagecodecs\nimport tifffile\nfrom scipy.ndimage import label\n\nprint(f\"imagecodecs version: {imagecodecs.__version__}\")\nprint(f\"tifffile version: {tifffile.__version__}\")\nprint(\"scipy loaded\")"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "import os\nimport gc\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport tifffile\nimport zipfile\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom scipy.ndimage import label\nfrom torch.amp import autocast\n\nIS_KAGGLE = os.path.exists('/kaggle/input')\nprint(f\"Running on Kaggle: {IS_KAGGLE}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cuda.matmul.allow_tf32 = True"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Configuration\nif IS_KAGGLE:\n    INPUT_PATH = '/kaggle/input/vesuvius-challenge-surface-detection'\n    V2_MODEL_PATH = '/kaggle/input/vesuvius-v2v3-ensemble-models/v2_best.pt'\n    V3_MODEL_PATH = '/kaggle/input/vesuvius-v2v3-ensemble-models/v3_best.pt'\n    OUTPUT_PATH = '/kaggle/working'\nelse:\n    INPUT_PATH = '/home/ck/Desktop/vesuvius/data'\n    V2_MODEL_PATH = '/home/ck/Desktop/vesuvius/outputs/nnunet_v2/best_model.pt'\n    V3_MODEL_PATH = '/home/ck/Desktop/vesuvius/outputs/nnunet_v3/best_model.pt'\n    OUTPUT_PATH = '/home/ck/Desktop/vesuvius/outputs'\n\n# Best config from local experiments\nPATCH_SIZE = (80, 80, 80)\nOVERLAP = 0.5\nTHRESHOLD = 0.25\nMIN_SIZE = 200\nUSE_TTA = True\nV2_WEIGHT = 1.0\nV3_WEIGHT = 0.5\n\nprint(f\"Config: patch={PATCH_SIZE}, overlap={OVERLAP}, thresh={THRESHOLD}\")\nprint(f\"Weights: V2={V2_WEIGHT}, V3={V3_WEIGHT}, TTA={USE_TTA}\")"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Model Definition - ResidualEncoderUNetSmall\n\nclass ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, norm='instance'):\n        super().__init__()\n        padding = kernel_size // 2\n        self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride=stride, padding=padding, bias=norm != 'batch')\n        if norm == 'instance':\n            self.norm = nn.InstanceNorm3d(out_channels, affine=True)\n        elif norm == 'batch':\n            self.norm = nn.BatchNorm3d(out_channels)\n        else:\n            self.norm = nn.Identity()\n        self.act = nn.LeakyReLU(0.01, inplace=True)\n\n    def forward(self, x):\n        return self.act(self.norm(self.conv(x)))\n\n\nclass ResidualBlock(nn.Module):\n    def __init__(self, channels, kernel_size=3, norm='instance'):\n        super().__init__()\n        self.conv1 = ConvBlock(channels, channels, kernel_size, norm=norm)\n        self.conv2 = nn.Sequential(\n            nn.Conv3d(channels, channels, kernel_size, padding=kernel_size // 2, bias=norm != 'batch'),\n            nn.InstanceNorm3d(channels, affine=True) if norm == 'instance' else nn.BatchNorm3d(channels),\n        )\n        self.act = nn.LeakyReLU(0.01, inplace=True)\n\n    def forward(self, x):\n        residual = x\n        out = self.conv1(x)\n        out = self.conv2(out)\n        return self.act(out + residual)\n\n\nclass EncoderStage(nn.Module):\n    def __init__(self, in_channels, out_channels, num_blocks=2, stride=2, norm='instance'):\n        super().__init__()\n        self.downsample = ConvBlock(in_channels, out_channels, kernel_size=3, stride=stride, norm=norm)\n        self.blocks = nn.Sequential(*[ResidualBlock(out_channels, norm=norm) for _ in range(num_blocks)])\n\n    def forward(self, x):\n        return self.blocks(self.downsample(x))\n\n\nclass DecoderStage(nn.Module):\n    def __init__(self, in_channels, skip_channels, out_channels, num_blocks=2, norm='instance'):\n        super().__init__()\n        self.upsample = nn.ConvTranspose3d(in_channels, out_channels, kernel_size=2, stride=2)\n        self.conv = ConvBlock(out_channels + skip_channels, out_channels, norm=norm)\n        self.blocks = nn.Sequential(*[ResidualBlock(out_channels, norm=norm) for _ in range(num_blocks)])\n\n    def forward(self, x, skip):\n        x = self.upsample(x)\n        if x.shape[2:] != skip.shape[2:]:\n            x = F.interpolate(x, size=skip.shape[2:], mode='trilinear', align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        return self.blocks(self.conv(x))\n\n\nclass ResidualEncoderUNetSmall(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1, base_features=32, max_features=256, deep_supervision=True):\n        super().__init__()\n        self.deep_supervision = deep_supervision\n        self.enc1 = nn.Sequential(ConvBlock(in_channels, base_features), ResidualBlock(base_features))\n        self.enc2 = EncoderStage(base_features, base_features * 2)\n        self.enc3 = EncoderStage(base_features * 2, base_features * 4)\n        self.enc4 = EncoderStage(base_features * 4, min(base_features * 8, max_features))\n        bottleneck_ch = min(base_features * 8, max_features)\n        self.bottleneck = nn.Sequential(ResidualBlock(bottleneck_ch), ResidualBlock(bottleneck_ch))\n        self.dec4 = DecoderStage(bottleneck_ch, base_features * 4, base_features * 4)\n        self.dec3 = DecoderStage(base_features * 4, base_features * 2, base_features * 2)\n        self.dec2 = DecoderStage(base_features * 2, base_features, base_features)\n        self.output = nn.Conv3d(base_features, out_channels, 1)\n        if deep_supervision:\n            self.ds3 = nn.Conv3d(base_features * 4, out_channels, 1)\n            self.ds2 = nn.Conv3d(base_features * 2, out_channels, 1)\n\n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(e1)\n        e3 = self.enc3(e2)\n        e4 = self.enc4(e3)\n        b = self.bottleneck(e4)\n        d4 = self.dec4(b, e3)\n        d3 = self.dec3(d4, e2)\n        d2 = self.dec2(d3, e1)\n        out = self.output(d2)\n        if self.deep_supervision and self.training:\n            return out, [self.ds3(d4), self.ds2(d3)]\n        return out"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "def clear_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\ndef create_gaussian_weight(patch_size, sigma=0.125):\n    d, h, w = patch_size\n    z = np.linspace(-1, 1, d)\n    y = np.linspace(-1, 1, h)\n    x = np.linspace(-1, 1, w)\n    zz, yy, xx = np.meshgrid(z, y, x, indexing='ij')\n    weight = np.exp(-(zz**2 + yy**2 + xx**2) / (2 * sigma**2))\n    return (weight / weight.max()).astype(np.float32)\n\ndef sliding_window_inference(model, volume, patch_size, overlap, device, use_tta=False):\n    \"\"\"Sliding window inference with Gaussian blending and optional 8x flip TTA.\"\"\"\n    model.eval()\n    d, h, w = volume.shape\n    pd, ph, pw = patch_size\n    stride = [int(p * (1 - overlap)) for p in patch_size]\n\n    positions = []\n    for zi in range(0, max(1, d - pd + 1), stride[0]):\n        for yi in range(0, max(1, h - ph + 1), stride[1]):\n            for xi in range(0, max(1, w - pw + 1), stride[2]):\n                positions.append((zi, yi, xi))\n    positions.append((max(0, d - pd), max(0, h - ph), max(0, w - pw)))\n    positions = list(set(positions))\n\n    gauss = create_gaussian_weight(patch_size)\n    pred = np.zeros((d, h, w), dtype=np.float32)\n    ws = np.zeros((d, h, w), dtype=np.float32)\n\n    with torch.no_grad():\n        for zi, yi, xi in tqdm(positions, desc=\"Inference\", leave=False):\n            patch = volume[zi:zi+pd, yi:yi+ph, xi:xi+pw]\n            if patch.shape != (pd, ph, pw):\n                p = np.zeros((pd, ph, pw), dtype=np.float32)\n                p[:patch.shape[0], :patch.shape[1], :patch.shape[2]] = patch\n                patch = p\n            pt = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).to(device)\n\n            if use_tta:\n                tta_preds = []\n                for flip_d in [False, True]:\n                    for flip_h in [False, True]:\n                        for flip_w in [False, True]:\n                            aug = pt\n                            if flip_d: aug = torch.flip(aug, [2])\n                            if flip_h: aug = torch.flip(aug, [3])\n                            if flip_w: aug = torch.flip(aug, [4])\n                            with autocast('cuda'):\n                                out = model(aug)\n                                if isinstance(out, tuple): out = out[0]\n                                out = torch.sigmoid(out)\n                            if flip_w: out = torch.flip(out, [4])\n                            if flip_h: out = torch.flip(out, [3])\n                            if flip_d: out = torch.flip(out, [2])\n                            tta_preds.append(out)\n                output = torch.stack(tta_preds).mean(0)\n            else:\n                with autocast('cuda'):\n                    output = model(pt)\n                    if isinstance(output, tuple): output = output[0]\n                    output = torch.sigmoid(output)\n\n            output = output.cpu().numpy()[0, 0]\n            ad, ah, aw = min(pd, d-zi), min(ph, h-yi), min(pw, w-xi)\n            pred[zi:zi+ad, yi:yi+ah, xi:xi+aw] += output[:ad, :ah, :aw] * gauss[:ad, :ah, :aw]\n            ws[zi:zi+ad, yi:yi+ah, xi:xi+aw] += gauss[:ad, :ah, :aw]\n\n    return pred / (ws + 1e-8)\n\ndef remove_small_components(binary_mask, min_size=200):\n    \"\"\"Remove connected components smaller than min_size.\"\"\"\n    labeled, n = label(binary_mask)\n    if n == 0:\n        return binary_mask\n    sizes = np.bincount(labeled.ravel())\n    mask = np.isin(labeled, np.where(sizes >= min_size)[0])\n    return (binary_mask * mask).astype(np.uint8)\n\ndef load_model(ckpt_path, device):\n    \"\"\"Load model from checkpoint.\"\"\"\n    model = ResidualEncoderUNetSmall(\n        in_channels=1, out_channels=1,\n        base_features=32, max_features=256,\n        deep_supervision=True,\n    ).to(device)\n    ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)\n    state = {k.replace('_orig_mod.', ''): v for k, v in ckpt['model_state_dict'].items()}\n    model.load_state_dict(state, strict=False)\n    model.eval()\n    return model"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Find test files\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {device}\")\n\ntest_dir = Path(INPUT_PATH) / 'test_images'\ntest_files = sorted(test_dir.glob('*.tif'))\n\nprint(f\"Test files: {len(test_files)}\")\nfor f in test_files:\n    print(f\"  - {f.name}\")"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Run ensemble inference\noutput_dir = Path(OUTPUT_PATH)\noutput_dir.mkdir(parents=True, exist_ok=True)\n\npredictions = {}\n\nfor test_file in test_files:\n    fid = test_file.stem\n    print(f\"\\n{'='*60}\")\n    print(f\"Processing: {fid}\")\n    print(f\"{'='*60}\")\n\n    # Load volume\n    img = tifffile.imread(str(test_file)).astype(np.float32) / 255.0\n    print(f\"Shape: {img.shape}, mean: {img.mean():.3f}\")\n\n    weighted_probs = np.zeros(img.shape, dtype=np.float64)\n    total_weight = 0.0\n\n    # V2 model\n    print(f\"\\nLoading V2 model (weight={V2_WEIGHT})...\")\n    model_v2 = load_model(V2_MODEL_PATH, device)\n    print(f\"Running V2 inference (TTA={USE_TTA})...\")\n    probs_v2 = sliding_window_inference(model_v2, img, PATCH_SIZE, OVERLAP, device, use_tta=USE_TTA)\n    weighted_probs += probs_v2.astype(np.float64) * V2_WEIGHT\n    total_weight += V2_WEIGHT\n    del model_v2\n    clear_memory()\n\n    # V3 model\n    print(f\"\\nLoading V3 model (weight={V3_WEIGHT})...\")\n    model_v3 = load_model(V3_MODEL_PATH, device)\n    print(f\"Running V3 inference (TTA={USE_TTA})...\")\n    probs_v3 = sliding_window_inference(model_v3, img, PATCH_SIZE, OVERLAP, device, use_tta=USE_TTA)\n    weighted_probs += probs_v3.astype(np.float64) * V3_WEIGHT\n    total_weight += V3_WEIGHT\n    del model_v3\n    clear_memory()\n\n    # Ensemble average\n    avg_probs = (weighted_probs / total_weight).astype(np.float32)\n    print(f\"\\nEnsemble probs: mean={avg_probs.mean():.4f}, std={avg_probs.std():.4f}\")\n\n    # Binarize + post-process\n    pred = (avg_probs > THRESHOLD).astype(np.uint8)\n    if MIN_SIZE > 0:\n        pred = remove_small_components(pred, min_size=MIN_SIZE)\n\n    fg_pct = pred.sum() / pred.size * 100\n    print(f\"Prediction: fg={pred.sum():,} ({fg_pct:.1f}%)\")\n\n    # Save\n    out_path = output_dir / f\"{fid}.tif\"\n    tifffile.imwrite(str(out_path), pred)\n    print(f\"Saved: {out_path}\")\n    predictions[fid] = pred\n    clear_memory()"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Create submission ZIP\nzip_path = output_dir / 'submission.zip'\nprint(f\"Creating submission: {zip_path}\")\n\nwith zipfile.ZipFile(str(zip_path), 'w', zipfile.ZIP_DEFLATED) as zf:\n    for fid in predictions:\n        tif_path = output_dir / f\"{fid}.tif\"\n        zf.write(str(tif_path), f\"{fid}.tif\")\n        print(f\"  Added: {fid}.tif\")\n\nprint(f\"\\n{'='*60}\")\nprint(f\"SUBMISSION READY\")\nprint(f\"{'='*60}\")\nprint(f\"Predictions: {len(predictions)} volumes\")\nprint(f\"ZIP: {zip_path}\")\nprint(f\"Config: V2={V2_WEIGHT}, V3={V3_WEIGHT}, TTA={USE_TTA}, thresh={THRESHOLD}, rm_small={MIN_SIZE}\")"}], "metadata": {"kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"}, "language_info": {"name": "python", "version": "3.10.12"}}, "nbformat": 4, "nbformat_minor": 4}