{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14958059,"datasetId":9443008,"databundleVersionId":15828974},{"sourceType":"modelInstanceVersion","sourceId":4535,"databundleVersionId":6346563,"modelInstanceId":3327,"modelId":986},{"sourceType":"kernelVersion","sourceId":290917305}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Vesuvius Surface Detection - DINOv2 2.5D Inference\n\n**Training Notebook**: [DINOv2 2.5D Segmentation Training](https://www.kaggle.com/code/rockerritesh/dinov2-2-5d-training-v2)\n\n**Key Features:**\n- DINOv2-Large backbone (1024-dim, 300M params)\n- UPerNet-style decoder with multi-scale fusion\n- 2.5D input: 5 consecutive slices (center +/- 2)\n- Sliding window inference with 4x rotation TTA\n- Hysteresis thresholding + 3D closing + dust removal","metadata":{}},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"# Install packages (offline mode for Kaggle)\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n    \"$var\"/tifffile-2025.12.20-py3-none-any.whl \\\n    \"$var\"/imagecodecs-2026.1.1-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl \\\n    --no-index \\\n    --find-links \"$var\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport math\nfrom pathlib import Path\nfrom typing import List, Tuple, Optional\nfrom functools import partial\n\nimport numpy as np\nimport pandas as pd\nimport tifffile\nfrom tqdm import tqdm\nimport zipfile\nfrom matplotlib import pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast\n\nfrom scipy.ndimage import binary_closing, binary_propagation, generate_binary_structure\nfrom skimage.morphology import remove_small_objects\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(f\"PyTorch: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"class CONFIG:\n    # Paths\n    ROOT_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection\"\n    TEST_DIR = f\"{ROOT_DIR}/test_images\"\n    OUTPUT_DIR = \"/kaggle/working/submission_masks\"\n    ZIP_PATH = \"/kaggle/working/submission.zip\"\n\n    # DINOv2 Hugging Face path (for reference only - NOT used directly)\n    # Available at: /kaggle/input/dinov2/pytorch/large/1\n    # Contains: config.json, preprocessor_config.json, pytorch_model.bin, README.md\n    # \n    # NOTE: We use inline ViT architecture instead of transformers library because:\n    # - Training used torch.hub.load('facebookresearch/dinov2', 'dinov2_vitl14')\n    # - torch.hub uses different attribute names than transformers:\n    #   * torch.hub: backbone.patch_embed.proj, backbone.blocks, backbone.pos_embed\n    #   * transformers: embeddings.patch_embeddings.projection, encoder.layer\n    # - The trained weights are saved with torch.hub naming convention\n    # - So we must use matching architecture (defined inline below)\n    DINO_HF_PATH = \"/kaggle/input/dinov2/pytorch/large/1\"  # For reference only\n\n    # Model weights path (your trained model - contains FULL model: backbone + decoder)\n    WEIGHTS_PATH = \"/kaggle/input/dinov2-2-5-v1/dinov2_vesuvius_best.pth\"\n\n    # Model config (must match training)\n    ENCODER_DIM = 1024  # DINOv2-Large\n    NUM_SLICES = 15  # 2.5D: center +/- 2\n    IMG_SIZE = 392  # Must match training\n    NUM_CLASSES = 2\n    PATCH_SIZE = 14  # DINOv2 patch size\n\n    # Inference\n    OVERLAP = 0.5\n    BATCH_SIZE = 4\n    USE_TTA = True  # 4x rotation TTA\n    USE_AMP = True\n\n    # Post-processing\n    T_LOW = 0.50\n    T_HIGH = 0.95\n    Z_RADIUS = 1\n    XY_RADIUS = 0\n    MIN_SIZE = 100\n\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nconfig = CONFIG()\nos.makedirs(config.OUTPUT_DIR, exist_ok=True)\nprint(f\"Device: {config.DEVICE}\")\nprint(f\"DINOv2 HF path (reference): {config.DINO_HF_PATH}\")\nprint(f\"Model weights: {config.WEIGHTS_PATH}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(f\"{config.ROOT_DIR}/test.csv\")\ntest_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DINOv2 ViT Architecture (Local Loading - No Network Required)\n\nThis section defines the DINOv2 Vision Transformer architecture so we can load weights from a local file without needing torch.hub or network access.","metadata":{}},{"cell_type":"code","source":"# =====================================================\n# DINOv2 ViT Architecture (from facebookresearch/dinov2)\n# Adapted for local loading without network access\n# IMPORTANT: Uses standard Mlp (fc1/fc2), NOT SwiGLUFFN\n# =====================================================\n\ndef drop_path(x, drop_prob: float = 0., training: bool = False):\n    if drop_prob == 0. or not training:\n        return x\n    keep_prob = 1 - drop_prob\n    shape = (x.shape[0],) + (1,) * (x.ndim - 1)\n    random_tensor = x.new_empty(shape).bernoulli_(keep_prob)\n    if keep_prob > 0.0:\n        random_tensor.div_(keep_prob)\n    return x * random_tensor\n\n\nclass DropPath(nn.Module):\n    def __init__(self, drop_prob=None):\n        super().__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x):\n        return drop_path(x, self.drop_prob, self.training)\n\n\nclass Mlp(nn.Module):\n    \"\"\"Standard MLP with fc1/fc2 (matches torch.hub DINOv2).\"\"\"\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0., bias=True):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features, bias=bias)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features, bias=bias)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\n\nclass Attention(nn.Module):\n    def __init__(self, dim, num_heads=8, qkv_bias=False, attn_drop=0., proj_drop=0.):\n        super().__init__()\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = head_dim ** -0.5\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x):\n        B, N, C = x.shape\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv.unbind(0)\n        \n        # Use scaled_dot_product_attention for efficiency\n        x = F.scaled_dot_product_attention(q, k, v)\n        x = x.transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n\nclass LayerScale(nn.Module):\n    def __init__(self, dim, init_values=1e-5, inplace=False):\n        super().__init__()\n        self.inplace = inplace\n        self.gamma = nn.Parameter(init_values * torch.ones(dim))\n\n    def forward(self, x):\n        return x.mul_(self.gamma) if self.inplace else x * self.gamma\n\n\nclass Block(nn.Module):\n    def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, drop=0., attn_drop=0.,\n                 drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, init_values=None):\n        super().__init__()\n        self.norm1 = norm_layer(dim)\n        self.attn = Attention(dim, num_heads=num_heads, qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop)\n        self.ls1 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity()\n        self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n        self.ls2 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity()\n        self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def forward(self, x):\n        x = x + self.drop_path1(self.ls1(self.attn(self.norm1(x))))\n        x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n        return x\n\n\nclass PatchEmbed(nn.Module):\n    def __init__(self, img_size=224, patch_size=14, in_chans=3, embed_dim=768):\n        super().__init__()\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = (img_size // patch_size) ** 2\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        x = self.proj(x)\n        x = x.flatten(2).transpose(1, 2)\n        return x\n\n\nclass DinoVisionTransformer(nn.Module):\n    \"\"\"DINOv2 Vision Transformer (matches torch.hub dinov2_vitl14).\"\"\"\n    def __init__(\n        self,\n        img_size=518,\n        patch_size=14,\n        in_chans=3,\n        embed_dim=1024,\n        depth=24,\n        num_heads=16,\n        mlp_ratio=4.,\n        qkv_bias=True,\n        init_values=1e-5,\n        drop_path_rate=0.,\n        norm_layer=partial(nn.LayerNorm, eps=1e-6),\n    ):\n        super().__init__()\n        self.num_features = self.embed_dim = embed_dim\n        self.num_tokens = 1\n        self.n_blocks = depth\n        self.num_heads = num_heads\n        self.patch_size = patch_size\n\n        self.patch_embed = PatchEmbed(img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))\n        self.mask_token = nn.Parameter(torch.zeros(1, embed_dim))  # DINOv2 has mask_token\n\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]\n        self.blocks = nn.ModuleList([\n            Block(\n                dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias,\n                drop_path=dpr[i], norm_layer=norm_layer, init_values=init_values\n            )\n            for i in range(depth)\n        ])\n        self.norm = norm_layer(embed_dim)\n\n    def interpolate_pos_encoding(self, x, w, h):\n        npatch = x.shape[1] - 1\n        N = self.pos_embed.shape[1] - 1\n        if npatch == N and w == h:\n            return self.pos_embed\n        \n        class_pos_embed = self.pos_embed[:, 0]\n        patch_pos_embed = self.pos_embed[:, 1:]\n        \n        dim = x.shape[-1]\n        w0 = w // self.patch_size\n        h0 = h // self.patch_size\n        \n        sqrt_N = int(math.sqrt(N))\n        patch_pos_embed = patch_pos_embed.reshape(1, sqrt_N, sqrt_N, dim).permute(0, 3, 1, 2)\n        patch_pos_embed = F.interpolate(patch_pos_embed, size=(h0, w0), mode='bicubic', align_corners=False)\n        patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).reshape(1, -1, dim)\n        \n        return torch.cat((class_pos_embed.unsqueeze(0), patch_pos_embed), dim=1)\n\n    def forward(self, x):\n        B, _, W, H = x.shape\n        x = self.patch_embed(x)\n        \n        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n        \n        x = x + self.interpolate_pos_encoding(x, W, H)\n        \n        for blk in self.blocks:\n            x = blk(x)\n        \n        x = self.norm(x)\n        return x\n\n\ndef create_dinov2_vitl14():\n    \"\"\"Create DINOv2-Large (ViT-L/14) model matching torch.hub version.\"\"\"\n    return DinoVisionTransformer(\n        img_size=518,\n        patch_size=14,\n        in_chans=3,\n        embed_dim=1024,\n        depth=24,\n        num_heads=16,\n        mlp_ratio=4.,\n        qkv_bias=True,\n        init_values=1e-5,\n    )\n\n\nprint(\"DINOv2 ViT architecture defined (matching torch.hub dinov2_vitl14)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architecture","metadata":{}},{"cell_type":"code","source":"class PPM(nn.Module):\n    \"\"\"Pyramid Pooling Module for global context.\"\"\"\n    def __init__(self, in_dim: int, reduction_dim: int, bins: Tuple[int, ...] = (1, 2, 3, 6)):\n        super().__init__()\n        self.features = nn.ModuleList()\n        for bin_size in bins:\n            self.features.append(nn.Sequential(\n                nn.AdaptiveAvgPool2d(bin_size),\n                nn.Conv2d(in_dim, reduction_dim, kernel_size=1, bias=False),\n                nn.BatchNorm2d(reduction_dim),\n                nn.ReLU(inplace=True)\n            ))\n        self.out_conv = nn.Sequential(\n            nn.Conv2d(in_dim + reduction_dim * len(bins), in_dim, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(in_dim),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        h, w = x.shape[2:]\n        pyramids = [x]\n        for f in self.features:\n            pyramids.append(F.interpolate(f(x), size=(h, w), mode='bilinear', align_corners=False))\n        return self.out_conv(torch.cat(pyramids, dim=1))\n\n\nclass UPerNetDecoder(nn.Module):\n    \"\"\"UPerNet-style decoder with multi-scale fusion.\"\"\"\n    def __init__(\n        self, \n        encoder_dim: int = 1024,\n        fpn_dim: int = 256,\n        num_classes: int = 2\n    ):\n        super().__init__()\n        \n        # PPM on deepest features\n        self.ppm = PPM(encoder_dim, fpn_dim // 4)\n        \n        # FPN lateral connections (4 scales from DINOv2 layers)\n        self.fpn_in = nn.ModuleList([\n            nn.Conv2d(encoder_dim, fpn_dim, kernel_size=1) for _ in range(4)\n        ])\n        \n        self.fpn_out = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(fpn_dim, fpn_dim, kernel_size=3, padding=1, bias=False),\n                nn.BatchNorm2d(fpn_dim),\n                nn.ReLU(inplace=True)\n            ) for _ in range(4)\n        ])\n        \n        # Fusion\n        self.fusion = nn.Sequential(\n            nn.Conv2d(fpn_dim * 4, fpn_dim, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(fpn_dim),\n            nn.ReLU(inplace=True)\n        )\n        \n        # Segmentation head\n        self.seg_head = nn.Sequential(\n            nn.Conv2d(fpn_dim, fpn_dim // 2, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(fpn_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(fpn_dim // 2, num_classes, kernel_size=1)\n        )\n    \n    def forward(self, features: List[torch.Tensor], target_size: Tuple[int, int]) -> torch.Tensor:\n        # Apply PPM to deepest feature\n        f4 = self.ppm(features[-1])\n        \n        # FPN top-down pathway\n        fpn_features = []\n        prev = None\n        for i in range(3, -1, -1):\n            f = self.fpn_in[i](features[i] if i < 3 else f4)\n            if prev is not None:\n                f = f + F.interpolate(prev, size=f.shape[2:], mode='bilinear', align_corners=False)\n            f = self.fpn_out[i](f)\n            fpn_features.insert(0, f)\n            prev = f\n        \n        # Upsample all to same size and concatenate\n        target_h, target_w = fpn_features[0].shape[2:]\n        upsampled = []\n        for f in fpn_features:\n            upsampled.append(F.interpolate(f, size=(target_h, target_w), mode='bilinear', align_corners=False))\n        \n        fused = self.fusion(torch.cat(upsampled, dim=1))\n        fused = F.interpolate(fused, size=target_size, mode='bilinear', align_corners=False)\n        \n        return self.seg_head(fused)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DINOv2Segmenter(nn.Module):\n    \"\"\"\n    DINOv2-Large + UPerNet for 2.5D segmentation.\n    \n    Input: (B, num_slices, H, W) - stacked CT slices\n    Output: (B, num_classes, H, W) - segmentation logits for center slice\n    \"\"\"\n    def __init__(\n        self,\n        num_slices: int = 5,\n        num_classes: int = 2,\n        encoder_dim: int = 1024,\n    ):\n        super().__init__()\n        self.num_slices = num_slices\n        self.encoder_dim = encoder_dim\n        self.patch_size = 14  # DINOv2 uses 14x14 patches\n        \n        # Create DINOv2-Large backbone (local, no network)\n        self.backbone = create_dinov2_vitl14()\n        \n        # Modify patch embedding to accept num_slices channels instead of 3\n        original_patch_embed = self.backbone.patch_embed.proj\n        new_patch_embed = nn.Conv2d(\n            num_slices, \n            encoder_dim, \n            kernel_size=original_patch_embed.kernel_size,\n            stride=original_patch_embed.stride,\n            padding=original_patch_embed.padding\n        )\n        self.backbone.patch_embed.proj = new_patch_embed\n        \n        # Store original position embedding for interpolation\n        self.register_buffer('orig_pos_embed', self.backbone.pos_embed.clone())\n        \n        # Feature extraction layers (DINOv2-Large has 24 layers)\n        self.feature_layers = [5, 11, 17, 23]\n        \n        # Decoder\n        self.decoder = UPerNetDecoder(\n            encoder_dim=encoder_dim,\n            fpn_dim=256,\n            num_classes=num_classes\n        )\n    \n    def interpolate_pos_embed(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:\n        \"\"\"Interpolate position embeddings to match input size.\"\"\"\n        npatch = h * w\n        N = self.orig_pos_embed.shape[1] - 1\n        \n        if npatch == N:\n            return self.orig_pos_embed\n        \n        cls_pos = self.orig_pos_embed[:, :1]\n        patch_pos = self.orig_pos_embed[:, 1:]\n        \n        orig_size = int(math.sqrt(N))\n        patch_pos = patch_pos.reshape(1, orig_size, orig_size, -1).permute(0, 3, 1, 2)\n        patch_pos = F.interpolate(patch_pos, size=(h, w), mode='bicubic', align_corners=False)\n        patch_pos = patch_pos.permute(0, 2, 3, 1).reshape(1, h * w, -1)\n        \n        return torch.cat([cls_pos, patch_pos], dim=1)\n    \n    def extract_features(self, x: torch.Tensor) -> List[torch.Tensor]:\n        \"\"\"Extract multi-scale features from DINOv2.\"\"\"\n        B, C, H, W = x.shape\n        \n        x = self.backbone.patch_embed(x)\n        h = H // self.patch_size\n        w = W // self.patch_size\n        \n        cls_token = self.backbone.cls_token.expand(B, -1, -1)\n        x = torch.cat([cls_token, x], dim=1)\n        \n        pos_embed = self.interpolate_pos_embed(x, h, w)\n        x = x + pos_embed\n        \n        features = []\n        for i, block in enumerate(self.backbone.blocks):\n            x = block(x)\n            if i in self.feature_layers:\n                feat = x[:, 1:, :]\n                feat = feat.permute(0, 2, 1).reshape(B, self.encoder_dim, h, w)\n                features.append(feat)\n        \n        return features\n    \n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        H, W = x.shape[2:]\n        features = self.extract_features(x)\n        logits = self.decoder(features, (H, W))\n        return logits\n\n\nprint(\"Model architecture defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Model","metadata":{}},{"cell_type":"code","source":"def get_model(weights_path: str, config) -> nn.Module:\n    \"\"\"Load model with trained weights (no network required).\"\"\"\n    # Create model\n    model = DINOv2Segmenter(\n        num_slices=config.NUM_SLICES,\n        num_classes=config.NUM_CLASSES,\n        encoder_dim=config.ENCODER_DIM,\n    )\n    \n    # Load trained weights (includes both backbone and decoder)\n    print(f\"Loading trained weights from: {weights_path}\")\n    state_dict = torch.load(weights_path, map_location='cpu')\n    model.load_state_dict(state_dict)\n    \n    model = model.to(config.DEVICE)\n    model.eval()\n    \n    print(f\"Model loaded: {model.decoder.seg_head[-1].out_channels} classes\")\n    print(f\"Parameters: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M\")\n    \n    return model\n\n\nmodel = get_model(config.WEIGHTS_PATH, config)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Sliding Window Inference","metadata":{}},{"cell_type":"code","source":"def _get_gaussian_kernel(size: int, sigma: float = None) -> np.ndarray:\n    \"\"\"Generate 2D Gaussian kernel for blending.\"\"\"\n    if sigma is None:\n        sigma = size / 4\n    x = np.arange(size) - size // 2\n    gauss_1d = np.exp(-x**2 / (2 * sigma**2))\n    gauss_2d = np.outer(gauss_1d, gauss_1d)\n    return gauss_2d / gauss_2d.max()\n\n\ndef sliding_window_inference_2d5(\n    model: nn.Module,\n    volume: np.ndarray,\n    roi_size: int = 392,\n    num_slices: int = 5,\n    overlap: float = 0.5,\n    batch_size: int = 4,\n    device: torch.device = torch.device('cuda'),\n    use_amp: bool = True\n) -> np.ndarray:\n    \"\"\"\n    2.5D sliding window inference.\n    \n    Args:\n        model: Trained model\n        volume: 3D CT volume (D, H, W)\n        roi_size: Patch size\n        num_slices: Number of slices for 2.5D\n        overlap: Overlap ratio\n        batch_size: Inference batch size\n        device: torch device\n        use_amp: Use mixed precision\n    \n    Returns:\n        probs: 3D probability map (D, H, W)\n    \"\"\"\n    model.eval()\n    D, H, W = volume.shape\n    half_slices = num_slices // 2\n    \n    # Pad volume if needed\n    pad_h = (roi_size - H % roi_size) % roi_size\n    pad_w = (roi_size - W % roi_size) % roi_size\n    \n    volume_padded = np.pad(volume, ((0, 0), (0, pad_h), (0, pad_w)), mode='reflect')\n    _, H_pad, W_pad = volume_padded.shape\n    \n    # Normalize to [0, 1]\n    volume_padded = volume_padded.astype(np.float32) / 255.0\n    \n    # Output arrays\n    probs = np.zeros((D, H_pad, W_pad), dtype=np.float32)\n    counts = np.zeros((D, H_pad, W_pad), dtype=np.float32)\n    \n    # Gaussian weight for blending\n    gaussian = _get_gaussian_kernel(roi_size)\n    \n    # Stride\n    stride = int(roi_size * (1 - overlap))\n    \n    # Generate all patches\n    patches = []\n    coords = []\n    \n    for z in range(half_slices, D - half_slices):\n        for y in range(0, H_pad - roi_size + 1, stride):\n            for x in range(0, W_pad - roi_size + 1, stride):\n                slices = volume_padded[z - half_slices : z + half_slices + 1, y:y+roi_size, x:x+roi_size]\n                patches.append(slices)\n                coords.append((z, y, x))\n    \n    # Batch inference\n    with torch.no_grad():\n        for i in tqdm(range(0, len(patches), batch_size), desc=f\"Total patches {len(patches)}\"):\n            batch = np.stack(patches[i:i+batch_size], axis=0)\n            batch_tensor = torch.from_numpy(batch).float().to(device)\n            \n            with autocast(enabled=use_amp):\n                logits = model(batch_tensor)\n                batch_probs = F.softmax(logits, dim=1)[:, 1].cpu().numpy()\n            \n            for j, (z, y, x) in enumerate(coords[i:i+batch_size]):\n                probs[z, y:y+roi_size, x:x+roi_size] += batch_probs[j] * gaussian\n                counts[z, y:y+roi_size, x:x+roi_size] += gaussian\n    \n    # Average overlapping regions\n    probs = np.divide(probs, counts, where=counts > 0)\n    \n    # Remove padding\n    probs = probs[:, :H, :W]\n    \n    return probs\n\n\nprint(\"Sliding window inference defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Rotation TTA","metadata":{}},{"cell_type":"code","source":"def rot90_volume(vol: np.ndarray, k: int) -> np.ndarray:\n    \"\"\"Rotate volume k times 90 degrees clockwise in HW plane.\"\"\"\n    return np.rot90(vol, k=-k, axes=(1, 2))\n\n\ndef unrot90_volume(vol: np.ndarray, k: int) -> np.ndarray:\n    \"\"\"Undo rotation.\"\"\"\n    return rot90_volume(vol, (4 - k) % 4)\n\n\ndef predict_with_tta(\n    model: nn.Module,\n    volume: np.ndarray,\n    roi_size: int = 392,\n    num_slices: int = 5,\n    overlap: float = 0.5,\n    batch_size: int = 4,\n    device: torch.device = torch.device('cuda'),\n    use_amp: bool = True\n) -> np.ndarray:\n    \"\"\"\n    Predict with 4x rotation TTA (0, 90, 180, 270 degrees).\n    \n    Args:\n        model: Trained model\n        volume: 3D CT volume (D, H, W)\n        roi_size: Patch size\n        num_slices: Number of slices for 2.5D\n        overlap: Overlap ratio\n        batch_size: Inference batch size\n        device: torch device\n        use_amp: Use mixed precision\n    \n    Returns:\n        probs: Averaged 3D probability map (D, H, W)\n    \"\"\"\n    probs_accum = []\n    \n    for k in range(4):\n        print(f\"TTA rotation {k * 90} degrees...\")\n        vol_rot = rot90_volume(volume, k)\n        \n        probs = sliding_window_inference_2d5(\n            model, vol_rot,\n            roi_size=roi_size,\n            num_slices=num_slices,\n            overlap=overlap,\n            batch_size=batch_size,\n            device=device,\n            use_amp=use_amp\n        )\n        \n        # Rotate back\n        probs = unrot90_volume(probs, k)\n        probs_accum.append(probs)\n    \n    return np.mean(probs_accum, axis=0)\n\n\nprint(\"TTA functions defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Post-Processing","metadata":{}},{"cell_type":"code","source":"def build_anisotropic_struct(z_radius: int, xy_radius: int) -> Optional[np.ndarray]:\n    \"\"\"Build anisotropic structuring element for 3D closing.\"\"\"\n    if z_radius == 0 and xy_radius == 0:\n        return None\n    \n    depth = 2 * z_radius + 1\n    size = 2 * xy_radius + 1 if xy_radius > 0 else 1\n    \n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz = z_radius\n    \n    for dz in range(-z_radius, z_radius + 1):\n        if xy_radius > 0:\n            cy = cx = xy_radius\n            for dy in range(-xy_radius, xy_radius + 1):\n                for dx in range(-xy_radius, xy_radius + 1):\n                    if dy * dy + dx * dx <= xy_radius * xy_radius:\n                        struct[cz + dz, cy + dy, cx + dx] = True\n        else:\n            struct[cz + dz, 0, 0] = True\n    \n    return struct\n\n\ndef postprocess(\n    probs: np.ndarray,\n    t_low: float = 0.50,\n    t_high: float = 0.90,\n    z_radius: int = 1,\n    xy_radius: int = 0,\n    min_size: int = 100\n) -> np.ndarray:\n    \"\"\"\n    Post-processing pipeline:\n    1. Hysteresis thresholding\n    2. Anisotropic 3D closing\n    3. Dust removal\n    \n    Args:\n        probs: 3D probability map (D, H, W)\n        t_low: Low threshold for hysteresis\n        t_high: High threshold for hysteresis\n        z_radius: Z radius for anisotropic closing\n        xy_radius: XY radius for anisotropic closing\n        min_size: Minimum component size to keep\n    \n    Returns:\n        Binary mask (D, H, W) uint8\n    \"\"\"\n    # 1. Hysteresis thresholding\n    strong = probs >= t_high\n    weak = probs >= t_low\n    \n    if not strong.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n    \n    struct_hyst = generate_binary_structure(3, 3)  # 26-connectivity\n    mask = binary_propagation(strong, mask=weak, structure=struct_hyst)\n    \n    if not mask.any():\n        return np.zeros_like(probs, dtype=np.uint8)\n    \n    # 2. Anisotropic 3D closing\n    if z_radius > 0 or xy_radius > 0:\n        struct_close = build_anisotropic_struct(z_radius, xy_radius)\n        if struct_close is not None:\n            mask = binary_closing(mask, structure=struct_close)\n    \n    # 3. Dust removal\n    if min_size > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=min_size)\n    \n    return mask.astype(np.uint8)\n\n\nprint(\"Post-processing functions defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualization","metadata":{}},{"cell_type":"code","source":"def visualize_prediction(\n    volume: np.ndarray,\n    probs: np.ndarray,\n    mask: np.ndarray,\n    slice_indices: List[int] = None,\n    figsize: Tuple[int, int] = (16, 12)\n):\n    \"\"\"\n    Visualize prediction results.\n    \n    Args:\n        volume: Original 3D volume (D, H, W)\n        probs: Probability map (D, H, W)\n        mask: Binary mask (D, H, W)\n        slice_indices: List of z-indices to visualize (default: evenly spaced)\n        figsize: Figure size\n    \"\"\"\n    D, H, W = volume.shape\n    \n    # Default: visualize 4 evenly spaced slices\n    if slice_indices is None:\n        slice_indices = [D // 5, 2 * D // 5, 3 * D // 5, 4 * D // 5]\n    \n    n_slices = len(slice_indices)\n    fig, axes = plt.subplots(3, n_slices, figsize=figsize)\n    \n    for i, z in enumerate(slice_indices):\n        # Original volume slice\n        axes[0, i].imshow(volume[z], cmap='gray')\n        axes[0, i].set_title(f'Input (z={z})')\n        axes[0, i].axis('off')\n        \n        # Probability map\n        axes[1, i].imshow(probs[z], cmap='hot', vmin=0, vmax=1)\n        axes[1, i].set_title(f'Probability (z={z})')\n        axes[1, i].axis('off')\n        \n        # Binary mask overlaid on original\n        axes[2, i].imshow(volume[z], cmap='gray')\n        axes[2, i].imshow(mask[z], cmap='Reds', alpha=0.5 * (mask[z] > 0))\n        axes[2, i].set_title(f'Prediction (z={z})')\n        axes[2, i].axis('off')\n    \n    axes[0, 0].set_ylabel('Input', fontsize=12)\n    axes[1, 0].set_ylabel('Probability', fontsize=12)\n    axes[2, 0].set_ylabel('Prediction', fontsize=12)\n    \n    plt.tight_layout()\n    plt.show()\n\n\ndef visualize_3d_projection(\n    mask: np.ndarray,\n    probs: np.ndarray = None,\n    figsize: Tuple[int, int] = (12, 4)\n):\n    \"\"\"\n    Visualize 3D mask using maximum intensity projections.\n    \n    Args:\n        mask: Binary mask (D, H, W)\n        probs: Optional probability map (D, H, W)\n        figsize: Figure size\n    \"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=figsize)\n    \n    # XY projection (top view)\n    xy_proj = mask.max(axis=0)\n    axes[0].imshow(xy_proj, cmap='gray')\n    axes[0].set_title('XY Projection (Top View)')\n    axes[0].axis('off')\n    \n    # XZ projection (front view)\n    xz_proj = mask.max(axis=1)\n    axes[1].imshow(xz_proj, cmap='gray', aspect='auto')\n    axes[1].set_title('XZ Projection (Front View)')\n    axes[1].axis('off')\n    \n    # YZ projection (side view)\n    yz_proj = mask.max(axis=2)\n    axes[2].imshow(yz_proj.T, cmap='gray', aspect='auto')\n    axes[2].set_title('YZ Projection (Side View)')\n    axes[2].axis('off')\n    \n    plt.suptitle(f'3D Mask Projections - Foreground voxels: {mask.sum():,}', fontsize=14)\n    plt.tight_layout()\n    plt.show()\n\n\nprint(\"Visualization functions defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prediction Functions","metadata":{}},{"cell_type":"code","source":"def load_volume(path: str) -> np.ndarray:\n    \"\"\"Load volume and ensure uint8 format.\"\"\"\n    vol = tifffile.imread(path)\n    if vol.dtype != np.uint8:\n        vol = ((vol - vol.min()) / (vol.max() - vol.min() + 1e-8) * 255).astype(np.uint8)\n    return vol\n\n\ndef predict_volume(\n    model: nn.Module,\n    volume: np.ndarray,\n    config,\n    save_probs_path: str = None\n) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"\n    Full prediction pipeline for a single volume.\n    \n    Args:\n        model: Trained model\n        volume: 3D CT volume (D, H, W)\n        config: Configuration object\n        save_probs_path: Optional path to save probability map\n    \n    Returns:\n        probs: Probability map (D, H, W)\n        mask: Binary mask (D, H, W) uint8\n    \"\"\"\n    # Predict with or without TTA\n    if config.USE_TTA:\n        probs = predict_with_tta(\n            model, volume,\n            roi_size=config.IMG_SIZE,\n            num_slices=config.NUM_SLICES,\n            overlap=config.OVERLAP,\n            batch_size=config.BATCH_SIZE,\n            device=config.DEVICE,\n            use_amp=config.USE_AMP\n        )\n    else:\n        probs = sliding_window_inference_2d5(\n            model, volume,\n            roi_size=config.IMG_SIZE,\n            num_slices=config.NUM_SLICES,\n            overlap=config.OVERLAP,\n            batch_size=config.BATCH_SIZE,\n            device=config.DEVICE,\n            use_amp=config.USE_AMP\n        )\n    \n    # Save probability map if requested\n    if save_probs_path is not None:\n        np.save(save_probs_path, probs)\n    \n    # Post-process\n    mask = postprocess(\n        probs,\n        t_low=config.T_LOW,\n        t_high=config.T_HIGH,\n        z_radius=config.Z_RADIUS,\n        xy_radius=config.XY_RADIUS,\n        min_size=config.MIN_SIZE\n    )\n    \n    return probs, mask\n\n\nprint(\"Prediction functions defined\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test Mode (Visualization)","metadata":{}},{"cell_type":"code","source":"# Set to True to run on training data for visualization\ntesting = False","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if testing:\n    # Use training data for visualization\n    test_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/test_images\"\n    test_df = pd.read_csv(f\"{config.ROOT_DIR}/train.csv\")\n    \n    # Select a few sample IDs for testing\n    test_ids = {1407735}\n    test_df = (\n        test_df\n        .loc[test_df[\"id\"].isin(test_ids)]\n        .reset_index(drop=True)\n    )\n    print(f\"Testing on {len(test_df)} samples\")\n    print(test_df)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prediction and Submission","metadata":{}},{"cell_type":"code","source":"# Create submission\nwith zipfile.ZipFile(config.ZIP_PATH, \"w\", compression=zipfile.ZIP_DEFLATED) as z:\n    for idx, row in test_df.iterrows():\n        image_id = row[\"id\"]\n        tif_path = f\"{config.TEST_DIR}/{image_id}.tif\"\n        \n        print(f\"\\n{'='*60}\")\n        print(f\"Processing {image_id} ({idx + 1}/{len(test_df)})...\")\n        print(f\"{'='*60}\")\n        \n        # Load volume\n        volume = load_volume(tif_path)\n        print(f\"Volume shape: {volume.shape}\")\n        \n        # Predict (save probs if testing)\n        probs_path = f\"{image_id}.npy\" if testing else None\n        probs, mask = predict_volume(model, volume, config, save_probs_path=probs_path)\n        print(f\"Mask shape: {mask.shape}, foreground voxels: {mask.sum():,}\")\n        \n        # Visualize if testing\n        if testing:\n            print(\"\\nVisualization:\")\n            visualize_prediction(volume, probs, mask)\n            visualize_3d_projection(mask, probs)\n        \n        # Save to zip\n        out_path = f\"{config.OUTPUT_DIR}/{image_id}.tif\"\n        tifffile.imwrite(out_path, mask.astype(np.uint8))\n        z.write(out_path, arcname=f\"{image_id}.tif\")\n        os.remove(out_path)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"Submission ZIP: {config.ZIP_PATH}\")\nprint(f\"{'='*60}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify submission\nwith zipfile.ZipFile(config.ZIP_PATH, 'r') as z:\n    print(f\"Files in submission: {len(z.namelist())}\")\n    for name in z.namelist()[:5]:\n        print(f\"  - {name}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}