{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"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":14772958,"datasetId":9443008,"databundleVersionId":15625228},{"sourceType":"datasetVersion","sourceId":14754877,"datasetId":9411017,"databundleVersionId":15605320},{"sourceType":"modelInstanceVersion","sourceId":732880,"databundleVersionId":15477237,"modelInstanceId":516822},{"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 - Ensemble Inference\n\n**Ensemble of two models with 0.5 weight each:**\n\n| Model | Architecture | Input | LB Score |\n|-------|-------------|-------|----------|\n| DINOv2 2.5D | DINOv2-Large + UPerNet | 2.5D (5 slices) | ~0.40 |\n| TransUNet 3D | SEResNeXt50 + TransUNet | 3D (160³) | ~0.545 |\n\n**Ensemble Strategy:**\n- Both models output probability maps\n- Simple averaging: `probs_ensemble = 0.5 * probs_dinov2 + 0.5 * probs_transunet`\n- Post-processing applied to ensemble probabilities","metadata":{}},{"cell_type":"markdown","source":"# Setup and Imports","metadata":{}},{"cell_type":"code","source":"# Install packages (offline mode for Kaggle)\nfrom IPython.display import clear_output\n\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    \"$var\"/keras_nightly-*.whl \\\n    \"$var\"/medicai-*.whl \\\n    --no-index \\\n    --find-links \"$var\"\n\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:09.076905Z","iopub.execute_input":"2026-02-10T01:22:09.077139Z","iopub.status.idle":"2026-02-10T01:22:17.347804Z","shell.execute_reply.started":"2026-02-10T01:22:09.077118Z","shell.execute_reply":"2026-02-10T01:22:17.347101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# ============ DUAL GPU SETUP ============\n# GPU 0: DINOv2 (PyTorch)\n# GPU 1: TransUNet (JAX/Keras)\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"  # Make both GPUs visible\nos.environ[\"JAX_PLATFORMS\"] = \"cuda\"\n\n# Configure JAX to use GPU 1 (index 1)\nos.environ[\"JAX_DEFAULT_DEVICE\"] = \"gpu\"\n\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\n# PyTorch for DINOv2 (will use GPU 0)\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast\n\n# Keras/JAX for TransUNet (will use GPU 1)\nimport jax\nimport keras\nfrom medicai.transforms import Compose, NormalizeIntensity\nfrom medicai.models import TransUNet\nfrom medicai.utils.inference import SlidingWindowInference\n\n# Post-processing\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\"Keras: {keras.__version__}\")\nprint(f\"JAX devices: {jax.devices()}\")\nprint(f\"PyTorch CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"PyTorch GPU count: {torch.cuda.device_count()}\")\n    for i in range(torch.cuda.device_count()):\n        print(f\"  GPU {i}: {torch.cuda.get_device_name(i)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:17.349330Z","iopub.execute_input":"2026-02-10T01:22:17.349583Z","iopub.status.idle":"2026-02-10T01:22:36.917575Z","shell.execute_reply.started":"2026-02-10T01:22:17.349556Z","shell.execute_reply":"2026-02-10T01:22:36.916884Z"}},"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 Model (GPU 0) ============\n    DINOV2_WEIGHTS = \"/kaggle/input/datasets/rockerritesh/dinov2-2-5-v1/dinov2_vesuvius_best.pth\"\n    DINOV2_ENCODER_DIM = 1024\n    DINOV2_NUM_SLICES = 5\n    DINOV2_IMG_SIZE = 392\n    DINOV2_NUM_CLASSES = 2\n    DINOV2_OVERLAP = 0.5\n    DINOV2_BATCH_SIZE = 4\n    DINOV2_USE_TTA = False\n    DINOV2_USE_AMP = True\n    DINOV2_DEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    # ============ TransUNet Model (GPU 1 via JAX) ============\n    TRANSUNET_WEIGHTS = \"/kaggle/input/vsd-model/keras/transunet/3/transunet.seresnext50.160px.comboloss.weights.h5\"\n    TRANSUNET_INPUT_SHAPE = (160, 160, 160)\n    TRANSUNET_NUM_CLASSES = 3\n    TRANSUNET_OVERLAP = 0.42\n\n    # ============ Ensemble ============\n    WEIGHT_DINOV2 = 0.5\n    WEIGHT_TRANSUNET = 0.5\n\n    # ============ Post-processing ============\n    T_LOW = 0.40\n    T_HIGH = 0.85\n    Z_RADIUS = 2\n    XY_RADIUS = 1\n    MIN_SIZE = 100\n\n    # Legacy (for compatibility)\n    DEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\nconfig = CONFIG()\nos.makedirs(config.OUTPUT_DIR, exist_ok=True)\nprint(f\"DINOv2 Device (PyTorch): {config.DINOV2_DEVICE}\")\nprint(f\"TransUNet Device (JAX): GPU 1 (via JAX)\")\nprint(f\"Ensemble weights: DINOv2={config.WEIGHT_DINOV2}, TransUNet={config.WEIGHT_TRANSUNET}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:36.918800Z","iopub.execute_input":"2026-02-10T01:22:36.919308Z","iopub.status.idle":"2026-02-10T01:22:36.926317Z","shell.execute_reply.started":"2026-02-10T01:22:36.919285Z","shell.execute_reply":"2026-02-10T01:22:36.925474Z"}},"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\")\nprint(f\"Test samples: {len(test_df)}\")\ntest_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:36.927482Z","iopub.execute_input":"2026-02-10T01:22:36.927899Z","iopub.status.idle":"2026-02-10T01:22:36.978312Z","shell.execute_reply.started":"2026-02-10T01:22:36.927867Z","shell.execute_reply":"2026-02-10T01:22:36.977556Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model 1: DINOv2 2.5D Architecture","metadata":{}},{"cell_type":"code","source":"# =====================================================\n# DINOv2 ViT Architecture (from facebookresearch/dinov2)\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    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        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        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    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))\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        class_pos_embed = self.pos_embed[:, 0]\n        patch_pos_embed = self.pos_embed[:, 1:]\n        dim = x.shape[-1]\n        w0 = w // self.patch_size\n        h0 = h // self.patch_size\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        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        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n        x = x + self.interpolate_pos_encoding(x, W, H)\n        for blk in self.blocks:\n            x = blk(x)\n        x = self.norm(x)\n        return x\n\n\ndef create_dinov2_vitl14():\n    return DinoVisionTransformer(\n        img_size=518, patch_size=14, in_chans=3, embed_dim=1024,\n        depth=24, num_heads=16, mlp_ratio=4., qkv_bias=True, init_values=1e-5,\n    )\n\n\nprint(\"DINOv2 ViT architecture defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:36.979956Z","iopub.execute_input":"2026-02-10T01:22:36.980522Z","iopub.status.idle":"2026-02-10T01:22:37.003427Z","shell.execute_reply.started":"2026-02-10T01:22:36.980501Z","shell.execute_reply":"2026-02-10T01:22:37.002796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# UPerNet Decoder\n\nclass PPM(nn.Module):\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    def __init__(self, encoder_dim: int = 1024, fpn_dim: int = 256, num_classes: int = 2):\n        super().__init__()\n        self.ppm = PPM(encoder_dim, fpn_dim // 4)\n        self.fpn_in = nn.ModuleList([nn.Conv2d(encoder_dim, fpn_dim, kernel_size=1) for _ in range(4)])\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        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        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        f4 = self.ppm(features[-1])\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        target_h, target_w = fpn_features[0].shape[2:]\n        upsampled = [F.interpolate(f, size=(target_h, target_w), mode='bilinear', align_corners=False) for f in fpn_features]\n        fused = self.fusion(torch.cat(upsampled, dim=1))\n        fused = F.interpolate(fused, size=target_size, mode='bilinear', align_corners=False)\n        return self.seg_head(fused)\n\n\nclass DINOv2Segmenter(nn.Module):\n    def __init__(self, num_slices: int = 5, num_classes: int = 2, encoder_dim: int = 1024):\n        super().__init__()\n        self.num_slices = num_slices\n        self.encoder_dim = encoder_dim\n        self.patch_size = 14\n        self.backbone = create_dinov2_vitl14()\n        original_patch_embed = self.backbone.patch_embed.proj\n        new_patch_embed = nn.Conv2d(\n            num_slices, 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        self.register_buffer('orig_pos_embed', self.backbone.pos_embed.clone())\n        self.feature_layers = [5, 11, 17, 23]\n        self.decoder = UPerNetDecoder(encoder_dim=encoder_dim, fpn_dim=256, num_classes=num_classes)\n    \n    def interpolate_pos_embed(self, x: torch.Tensor, h: int, w: int) -> torch.Tensor:\n        npatch = h * w\n        N = self.orig_pos_embed.shape[1] - 1\n        if npatch == N:\n            return self.orig_pos_embed\n        cls_pos = self.orig_pos_embed[:, :1]\n        patch_pos = self.orig_pos_embed[:, 1:]\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        return torch.cat([cls_pos, patch_pos], dim=1)\n    \n    def extract_features(self, x: torch.Tensor) -> List[torch.Tensor]:\n        B, C, H, W = x.shape\n        x = self.backbone.patch_embed(x)\n        h = H // self.patch_size\n        w = W // self.patch_size\n        cls_token = self.backbone.cls_token.expand(B, -1, -1)\n        x = torch.cat([cls_token, x], dim=1)\n        pos_embed = self.interpolate_pos_embed(x, h, w)\n        x = x + pos_embed\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        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(\"DINOv2 Segmenter defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:37.004367Z","iopub.execute_input":"2026-02-10T01:22:37.004580Z","iopub.status.idle":"2026-02-10T01:22:37.023755Z","shell.execute_reply.started":"2026-02-10T01:22:37.004562Z","shell.execute_reply":"2026-02-10T01:22:37.023090Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model 2: TransUNet 3D","metadata":{}},{"cell_type":"code","source":"def get_transunet_model(config):\n    model = TransUNet(\n        input_shape=(*config.TRANSUNET_INPUT_SHAPE, 1),\n        encoder_name='seresnext50',\n        classifier_activation=None,\n        num_classes=config.TRANSUNET_NUM_CLASSES,\n    )\n    model.load_weights(config.TRANSUNET_WEIGHTS)\n    return model\n\n\nprint(\"TransUNet model function defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:37.024689Z","iopub.execute_input":"2026-02-10T01:22:37.024996Z","iopub.status.idle":"2026-02-10T01:22:37.038513Z","shell.execute_reply.started":"2026-02-10T01:22:37.024967Z","shell.execute_reply":"2026-02-10T01:22:37.037961Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Both Models","metadata":{}},{"cell_type":"code","source":"# Load DINOv2 model on GPU 0\ndef load_dinov2_model(config):\n    model = DINOv2Segmenter(\n        num_slices=config.DINOV2_NUM_SLICES,\n        num_classes=config.DINOV2_NUM_CLASSES,\n        encoder_dim=config.DINOV2_ENCODER_DIM,\n    )\n    print(f\"Loading DINOv2 weights from: {config.DINOV2_WEIGHTS}\")\n    state_dict = torch.load(config.DINOV2_WEIGHTS, map_location='cpu')\n    model.load_state_dict(state_dict)\n    model = model.to(config.DINOV2_DEVICE)  # GPU 0\n    model.eval()\n    print(f\"DINOv2 loaded on: {config.DINOV2_DEVICE}\")\n    print(f\"DINOv2 parameters: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M\")\n    return model\n\n\nprint(\"Loading DINOv2 on GPU 0...\")\ndinov2_model = load_dinov2_model(config)\n\n# Clear PyTorch cache before loading JAX model\ntorch.cuda.empty_cache()\n\n# Load TransUNet model on GPU 1 via JAX\nprint(f\"\\n{'='*50}\")\nprint(\"Loading TransUNet on GPU 1 (JAX)...\")\nprint(f\"JAX default backend: {jax.default_backend()}\")\nprint(f\"JAX devices available: {jax.devices()}\")\n\n# Force JAX to use GPU 1\nwith jax.default_device(jax.devices('gpu')[1] if len(jax.devices('gpu')) > 1 else jax.devices()[0]):\n    print(f\"Loading TransUNet weights from: {config.TRANSUNET_WEIGHTS}\")\n    transunet_model = get_transunet_model(config)\n    print(f\"TransUNet parameters: {transunet_model.count_params() / 1e6:.2f}M\")\n    \n    # Create sliding window inference for TransUNet\n    transunet_swi = SlidingWindowInference(\n        transunet_model,\n        num_classes=config.TRANSUNET_NUM_CLASSES,\n        roi_size=config.TRANSUNET_INPUT_SHAPE,\n        sw_batch_size=1,\n        mode='gaussian',\n        overlap=config.TRANSUNET_OVERLAP,\n    )\n\nprint(f\"\\n{'='*50}\")\nprint(\"Both models loaded successfully on separate GPUs!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:22:37.039390Z","iopub.execute_input":"2026-02-10T01:22:37.039901Z","iopub.status.idle":"2026-02-10T01:23:08.430376Z","shell.execute_reply.started":"2026-02-10T01:22:37.039881Z","shell.execute_reply":"2026-02-10T01:23:08.429723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DINOv2 Inference Functions","metadata":{}},{"cell_type":"code","source":"def _get_gaussian_kernel(size: int, sigma: float = None) -> np.ndarray:\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_dinov2(\n    model: nn.Module,\n    volume: np.ndarray,\n    roi_size: int,\n    num_slices: int,\n    overlap: float,\n    batch_size: int,\n    device: torch.device,\n    use_amp: bool\n) -> np.ndarray:\n    \"\"\"2.5D sliding window inference returning probability map.\"\"\"\n    model.eval()\n    D, H, W = volume.shape\n    half_slices = num_slices // 2\n    \n    pad_h = (roi_size - H % roi_size) % roi_size\n    pad_w = (roi_size - W % roi_size) % roi_size\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    volume_padded = volume_padded.astype(np.float32) / 255.0\n    \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    gaussian = _get_gaussian_kernel(roi_size)\n    stride = int(roi_size * (1 - overlap))\n    \n    patches = []\n    coords = []\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    with torch.no_grad():\n        for i in tqdm(range(0, len(patches), batch_size), desc=\"DINOv2 (GPU 0)\"):\n            batch = np.stack(patches[i:i+batch_size], axis=0)\n            batch_tensor = torch.from_numpy(batch).float().to(device)\n            with autocast(enabled=use_amp):\n                logits = model(batch_tensor)\n                batch_probs = F.softmax(logits, dim=1)[:, 1].cpu().numpy()\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    probs = np.divide(probs, counts, where=counts > 0)\n    probs = probs[:, :H, :W]\n    return probs\n\n\ndef rot90_volume(vol: np.ndarray, k: int) -> np.ndarray:\n    return np.rot90(vol, k=-k, axes=(1, 2))\n\n\ndef unrot90_volume(vol: np.ndarray, k: int) -> np.ndarray:\n    return rot90_volume(vol, (4 - k) % 4)\n\n\ndef predict_dinov2_with_tta(model, volume, config) -> np.ndarray:\n    \"\"\"DINOv2 prediction with 4x rotation TTA, returns probability map.\"\"\"\n    probs_accum = []\n    for k in range(4):\n        print(f\"  DINOv2 TTA rotation {k * 90}...\")\n        vol_rot = rot90_volume(volume, k)\n        probs = sliding_window_inference_dinov2(\n            model, vol_rot,\n            roi_size=config.DINOV2_IMG_SIZE,\n            num_slices=config.DINOV2_NUM_SLICES,\n            overlap=config.DINOV2_OVERLAP,\n            batch_size=config.DINOV2_BATCH_SIZE,\n            device=config.DINOV2_DEVICE,  # Use DINOV2_DEVICE (GPU 0)\n            use_amp=config.DINOV2_USE_AMP\n        )\n        probs = unrot90_volume(probs, k)\n        probs_accum.append(probs)\n    return np.mean(probs_accum, axis=0)\n\n\nprint(\"DINOv2 inference functions defined (GPU 0)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.431260Z","iopub.execute_input":"2026-02-10T01:23:08.431519Z","iopub.status.idle":"2026-02-10T01:23:08.445467Z","shell.execute_reply.started":"2026-02-10T01:23:08.431500Z","shell.execute_reply":"2026-02-10T01:23:08.444775Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TransUNet Inference Functions","metadata":{}},{"cell_type":"code","source":"def transunet_normalize(volume: np.ndarray) -> np.ndarray:\n    \"\"\"Simple z-score normalization on CPU (avoids GPU memory issues).\"\"\"\n    # Use numpy instead of TensorFlow/JAX for normalization to avoid GPU memory issues\n    vol = volume.astype(np.float32)\n    mask = vol != 0\n    if mask.any():\n        mean_val = vol[mask].mean()\n        std_val = vol[mask].std()\n        if std_val > 0:\n            vol = (vol - mean_val) / std_val\n    return vol\n\n\ndef predict_transunet_with_tta(swi, volume: np.ndarray) -> np.ndarray:\n    \"\"\"TransUNet prediction with TTA on GPU 1, returns probability map (foreground class).\"\"\"\n    # Normalize on CPU first\n    volume_norm = transunet_normalize(volume)\n    \n    # Prepare input: (1, D, H, W, 1)\n    inputs = volume_norm[None, ..., None]\n    \n    logits_list = []\n    \n    # Select GPU 1 for JAX operations\n    gpu1 = jax.devices('gpu')[1] if len(jax.devices('gpu')) > 1 else jax.devices()[0]\n    \n    with jax.default_device(gpu1):\n        # Original\n        print(\"  TransUNet TTA: original (GPU 1)\")\n        logits_list.append(swi(inputs))\n        \n        # Flips (spatial only)\n        for axis in [1, 2, 3]:\n            print(f\"  TransUNet TTA: flip axis {axis} (GPU 1)\")\n            img_f = np.flip(inputs, axis=axis).copy()  # .copy() ensures contiguous array\n            p = swi(img_f)\n            p = np.flip(p, axis=axis).copy()\n            logits_list.append(p)\n        \n        # Axial rotations (H, W)\n        for k in [1, 2, 3]:\n            print(f\"  TransUNet TTA: rot90 k={k} (GPU 1)\")\n            img_r = np.rot90(inputs, k=k, axes=(2, 3)).copy()\n            p = swi(img_r)\n            p = np.rot90(p, k=-k, axes=(2, 3)).copy()\n            logits_list.append(p)\n    \n    # Average logits and apply softmax (on CPU)\n    mean_logits = np.mean(logits_list, axis=0)  # (1, D, H, W, 3)\n    \n    # Convert to probabilities using softmax\n    # For 3-class model: class 0=background, class 1=foreground, class 2=unlabeled\n    from scipy.special import softmax\n    probs = softmax(mean_logits, axis=-1)  # (1, D, H, W, 3)\n    foreground_probs = probs[0, ..., 1]  # (D, H, W) - foreground class probability\n    \n    return foreground_probs\n\n\nprint(\"TransUNet inference functions defined (GPU 1)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.446439Z","iopub.execute_input":"2026-02-10T01:23:08.446804Z","iopub.status.idle":"2026-02-10T01:23:08.531640Z","shell.execute_reply.started":"2026-02-10T01:23:08.446784Z","shell.execute_reply":"2026-02-10T01:23:08.530906Z"}},"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    if z_radius == 0 and xy_radius == 0:\n        return None\n    depth = 2 * z_radius + 1\n    size = 2 * xy_radius + 1 if xy_radius > 0 else 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz = z_radius\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    return struct\n\n\ndef postprocess(\n    probs: np.ndarray,\n    t_low: float,\n    t_high: float,\n    z_radius: int,\n    xy_radius: int,\n    min_size: int\n) -> np.ndarray:\n    \"\"\"Post-processing: hysteresis + 3D closing + dust removal.\"\"\"\n    # 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)\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    # 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    # 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,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.532428Z","iopub.execute_input":"2026-02-10T01:23:08.532622Z","iopub.status.idle":"2026-02-10T01:23:08.544730Z","shell.execute_reply.started":"2026-02-10T01:23:08.532605Z","shell.execute_reply":"2026-02-10T01:23:08.544113Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ensemble Prediction","metadata":{}},{"cell_type":"code","source":"def load_volume(path: str) -> np.ndarray:\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 ensemble_predict(\n    dinov2_model,\n    transunet_swi,\n    volume: np.ndarray,\n    config\n) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"\n    Ensemble prediction combining DINOv2 (GPU 0) and TransUNet (GPU 1).\n    \n    Returns:\n        probs_ensemble: Combined probability map\n        mask: Binary mask after post-processing\n    \"\"\"\n    D, H, W = volume.shape\n    print(f\"Volume shape: {volume.shape}\")\n    \n    # ============ DINOv2 Prediction (GPU 0) ============\n    print(\"\\n[1/2] Running DINOv2 2.5D inference on GPU 0...\")\n    if config.DINOV2_USE_TTA:\n        probs_dinov2 = predict_dinov2_with_tta(dinov2_model, volume, config)\n    else:\n        probs_dinov2 = sliding_window_inference_dinov2(\n            dinov2_model, volume,\n            roi_size=config.DINOV2_IMG_SIZE,\n            num_slices=config.DINOV2_NUM_SLICES,\n            overlap=config.DINOV2_OVERLAP,\n            batch_size=config.DINOV2_BATCH_SIZE,\n            device=config.DINOV2_DEVICE,\n            use_amp=config.DINOV2_USE_AMP\n        )\n    print(f\"DINOv2 probs shape: {probs_dinov2.shape}, range: [{probs_dinov2.min():.3f}, {probs_dinov2.max():.3f}]\")\n    \n    # Clear PyTorch GPU cache after DINOv2 inference\n    torch.cuda.empty_cache()\n    \n    # ============ TransUNet Prediction (GPU 1) ============\n    print(\"\\n[2/2] Running TransUNet 3D inference on GPU 1...\")\n    volume_float = volume.astype(np.float32)\n    probs_transunet = predict_transunet_with_tta(transunet_swi, volume_float)\n    print(f\"TransUNet probs shape: {probs_transunet.shape}, range: [{probs_transunet.min():.3f}, {probs_transunet.max():.3f}]\")\n    \n    # ============ Ensemble ============\n    print(\"\\n[Ensemble] Combining predictions...\")\n    probs_ensemble = config.WEIGHT_DINOV2 * probs_dinov2 + config.WEIGHT_TRANSUNET * probs_transunet\n    print(f\"Ensemble probs range: [{probs_ensemble.min():.3f}, {probs_ensemble.max():.3f}]\")\n    \n    # Free memory\n    del probs_dinov2, probs_transunet\n    \n    # ============ Post-processing ============\n    print(\"\\n[Post-processing]...\")\n    mask = postprocess(\n        probs_ensemble,\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    print(f\"Mask shape: {mask.shape}, foreground voxels: {mask.sum():,}\")\n    \n    return probs_ensemble, mask\n\n\nprint(\"Ensemble prediction function defined (DINOv2 on GPU 0, TransUNet on GPU 1)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.545582Z","iopub.execute_input":"2026-02-10T01:23:08.545834Z","iopub.status.idle":"2026-02-10T01:23:08.560937Z","shell.execute_reply.started":"2026-02-10T01:23:08.545807Z","shell.execute_reply":"2026-02-10T01:23:08.560381Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization","metadata":{}},{"cell_type":"code","source":"def visualize_ensemble(\n    volume: np.ndarray,\n    probs_dinov2: np.ndarray,\n    probs_transunet: np.ndarray,\n    probs_ensemble: np.ndarray,\n    mask: np.ndarray,\n    slice_idx: int = None\n):\n    \"\"\"Visualize ensemble results.\"\"\"\n    D, H, W = volume.shape\n    if slice_idx is None:\n        slice_idx = D // 2\n    \n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    axes[0, 0].imshow(volume[slice_idx], cmap='gray')\n    axes[0, 0].set_title(f'Input (z={slice_idx})')\n    axes[0, 0].axis('off')\n    \n    axes[0, 1].imshow(probs_dinov2[slice_idx], cmap='hot', vmin=0, vmax=1)\n    axes[0, 1].set_title(f'DINOv2 Probs (w={config.WEIGHT_DINOV2})')\n    axes[0, 1].axis('off')\n    \n    axes[0, 2].imshow(probs_transunet[slice_idx], cmap='hot', vmin=0, vmax=1)\n    axes[0, 2].set_title(f'TransUNet Probs (w={config.WEIGHT_TRANSUNET})')\n    axes[0, 2].axis('off')\n    \n    axes[1, 0].imshow(probs_ensemble[slice_idx], cmap='hot', vmin=0, vmax=1)\n    axes[1, 0].set_title('Ensemble Probs')\n    axes[1, 0].axis('off')\n    \n    axes[1, 1].imshow(mask[slice_idx], cmap='gray')\n    axes[1, 1].set_title('Final Mask')\n    axes[1, 1].axis('off')\n    \n    axes[1, 2].imshow(volume[slice_idx], cmap='gray')\n    axes[1, 2].imshow(mask[slice_idx], cmap='Reds', alpha=0.5 * (mask[slice_idx] > 0))\n    axes[1, 2].set_title('Overlay')\n    axes[1, 2].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n\nprint(\"Visualization functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.561749Z","iopub.execute_input":"2026-02-10T01:23:08.561981Z","iopub.status.idle":"2026-02-10T01:23:08.576200Z","shell.execute_reply.started":"2026-02-10T01:23:08.561959Z","shell.execute_reply":"2026-02-10T01:23:08.575611Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test Mode (Optional)","metadata":{}},{"cell_type":"code","source":"# Set to True to run on training data for visualization\ntesting = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.578223Z","iopub.execute_input":"2026-02-10T01:23:08.578455Z","iopub.status.idle":"2026-02-10T01:23:08.590249Z","shell.execute_reply.started":"2026-02-10T01:23:08.578437Z","shell.execute_reply":"2026-02-10T01:23:08.589484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if testing:\n    test_dir = \"/kaggle/input/vesuvius-challenge-surface-detection/train_images\"\n    test_df = pd.read_csv(f\"{config.ROOT_DIR}/train.csv\")\n    test_ids = {1407735}  # Single sample for testing\n    test_df = test_df.loc[test_df[\"id\"].isin(test_ids)].reset_index(drop=True)\n    print(f\"Testing on {len(test_df)} samples\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.591071Z","iopub.execute_input":"2026-02-10T01:23:08.591241Z","iopub.status.idle":"2026-02-10T01:23:08.597830Z","shell.execute_reply.started":"2026-02-10T01:23:08.591226Z","shell.execute_reply":"2026-02-10T01:23:08.597104Z"}},"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{'='*70}\")\n        print(f\"Processing {image_id} ({idx + 1}/{len(test_df)})\")\n        print(f\"{'='*70}\")\n        \n        # Load volume\n        volume = load_volume(tif_path)\n        \n        # Ensemble prediction\n        probs_ensemble, mask = ensemble_predict(\n            dinov2_model,\n            transunet_swi,\n            volume,\n            config\n        )\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{'='*70}\")\nprint(f\"Submission ZIP: {config.ZIP_PATH}\")\nprint(f\"{'='*70}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T01:23:08.598601Z","iopub.execute_input":"2026-02-10T01:23:08.598827Z","iopub.status.idle":"2026-02-10T01:25:59.756607Z","shell.execute_reply.started":"2026-02-10T01:23:08.598804Z","shell.execute_reply":"2026-02-10T01:25:59.755984Z"}},"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,"execution":{"iopub.status.busy":"2026-02-10T01:25:59.757492Z","iopub.execute_input":"2026-02-10T01:25:59.757720Z","iopub.status.idle":"2026-02-10T01:25:59.763719Z","shell.execute_reply.started":"2026-02-10T01:25:59.757699Z","shell.execute_reply":"2026-02-10T01:25:59.763068Z"}},"outputs":[],"execution_count":null}]}