{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"gpu","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":726146,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":552759,"modelId":565313},{"sourceId":728126,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":554416,"modelId":566977},{"sourceId":730156,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":556146,"modelId":568710}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python3\n# -*- coding: utf-8 -*-\n\"\"\"\nHYBRID 2.5D x3 (ori1/ori2/ori3) + STACKING (Fusion UNet) — INFERENCE\nRespeta la estructura del script base (rutas, outputs, zip, tiled Hann, AMP, etc).\n\nPipeline:\n1) Carga 3 checkpoints base (HybridModel) entrenados en 3 orientaciones distintas.\n2) Para cada volumen test:\n   - Genera 3 volúmenes de logits2 (uno por orientación) con inferencia tiled + Hann.\n   - Reproyecta cada orientación a grilla canónica (Z,Y,X).\n3) Para cada slice z (plano XY canónico):\n   - Infiere máscara final con el modelo de stacking (checkpoint #4) usando tiles + Hann:\n        inputs por tile: x25d (Zwin,CROP,CROP) + logits3 (3,CROP,CROP).\n4) Escribe OUT_DIR/<id>.tif (multipage D páginas uint8 0/1) y ZIP_PATH con todos los tifs.\n\nRequisitos:\n  pip install tifffile pillow numpy torch tqdm\n\"\"\"\n\nimport os\nos.environ.setdefault(\"PYTORCH_CUDA_ALLOC_CONF\", \"expandable_segments:True\")\n\nimport math\nimport zipfile\nfrom pathlib import Path\nfrom typing import Dict, Any, Tuple, List, Optional\n\nimport numpy as np\nimport tifffile\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom tqdm.auto import tqdm\nfrom contextlib import nullcontext\n\n\n# ============================================================\n# CONFIG  (misma estructura del script base + 4 CKPTs)\n# ============================================================\n\n# --- Paths ---\nROOT_DIR = \"/kaggle/input/vesuvius-challenge-surface-detection\"\nTEST_DIR = f\"{ROOT_DIR}/test_images\"\n\nOUT_DIR  = \"/kaggle/working/submission_masks\"\nZIP_PATH = \"/kaggle/working/submission.zip\"\n\n# --- 4 checkpoints: 3 bases + stacking head ---\nCKPT_ORI1  = \"/kaggle/input/stacked/pytorch/default/1/best_1.pt\"\nCKPT_ORI2  = \"/kaggle/input/stacked/pytorch/default/1/best_2.pt\"\nCKPT_ORI3  = \"/kaggle/input/stacked/pytorch/default/1/best_3.pt\"\nCKPT_STACK = \"/kaggle/input/refiner-big/pytorch/default/1/best_post_loss.pt\"\n\n# Device principal (para DP)\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# MultiGPU\nUSE_MULTI_GPU = False\nGPU_IDS = [0, 1]  # ajustá según tu máquina\n\n# Inference\nUSE_AMP = True\nUSE_BF16 = True   # recomendado (si tu GPU lo soporta)\n\n# Usar pesos EMA si existen en el ckpt (ck[\"ema\"][\"shadow\"])\nPREFER_EMA_WEIGHTS = True\n\n# Threshold\nTHRESHOLD = 0.35  # -1 => intenta best_thr del ckpt STACK si existe; sino ori1; sino 0.5\n\n# Ventana 2.5D (debe coincidir con base models)\nCROP = 320\nZ_WINDOW = 9\nPATCH = 16\nOVERLAP_PATCH = False\nEMBED = 512\nLAYERS = 8\nHEADS = 8\nDROP = 0.05\nDROPPATH = 0.10\nBASE_CH = 48\nREFINE_DEPTH = 2\nREFINE_FEATURES = (80, 160, 320, 640)\nDETACH_STAGE1_FOR_STAGE2 = True\n\n# Sliding window 2D (H,W)\n# STRIDE=0 => default: CROP//2\nSTRIDE = 0\n\n# Batch de tiles por forward (para que DataParallel se note: >= 2 * #GPUs)\nBATCH_TILES = 8\n\n# Normalización (MATCH training: norm por patch Z×H×W) + clamp\nNORMALIZE = \"zscore\"   # \"none\" | \"zscore\" | \"minmax\"\nCLAMP_AFTER_NORM = 6.0 # x = x.clamp(-6,6)\n\n# Output mode:\n#   \"3d\"       => escribe multipage D,H,W (una página por z)\n#   \"2d_mean\"  => promedia probs a lo largo de z y escribe 2D\n#   \"2d_max\"   => max sobre z y escribe 2D\nOUTPUT_MODE = \"3d\"\n\n# Si querés contorno en lugar de máscara llena\nOUTLINE_ONLY = False\nOUTLINE_CONN8 = True\n\n# --- Stacking head config (debe coincidir con trainer) ---\nFUSION_BASE = 64\nFUSION_DROP = 0.05\nFUSION_USE_BND = False  # si lo entrenaste con boundary head, ponelo True\n\n\n# ============================================================\n# TIFF loader robusto (con memmap cuando se puede)\n# ============================================================\n\ndef load_tiff_volume_no_codec(path: str) -> np.ndarray:\n    \"\"\"\n    Return (D,H,W).\n    - Intenta tifffile.memmap -> evita RAM grande si es posible\n    - Fallback a tifffile.imread\n    - Fallback a PIL multipage si hay issues de codec/compresión\n    \"\"\"\n    try:\n        mm = tifffile.memmap(path)\n        vol = np.asarray(mm)\n        if vol.ndim == 2:\n            vol = vol[None, ...]\n        if vol.ndim != 3:\n            raise ValueError(f\"Expected 2D/3D tiff, got shape={vol.shape}\")\n        return vol\n    except Exception:\n        try:\n            vol = tifffile.imread(path)\n            if vol.ndim == 2:\n                vol = vol[None, ...]\n            if vol.ndim != 3:\n                raise ValueError(f\"Expected 2D/3D tiff, got shape={vol.shape}\")\n            return vol\n        except Exception:\n            img = Image.open(path)\n            slices = []\n            i = 0\n            while True:\n                try:\n                    img.seek(i)\n                    slices.append(np.array(img))\n                    i += 1\n                except EOFError:\n                    break\n            vol = np.stack(slices, axis=0)\n            return vol\n\ndef maybe_fix_dim_order(vol: np.ndarray) -> np.ndarray:\n    \"\"\"\n    Heurística por si viene (H,W,D).\n    Asume D suele ser el eje \"pequeño\".\n    \"\"\"\n    if vol.ndim != 3:\n        return vol\n    D, H, W = vol.shape\n    # heurística simple\n    if D > 1024 and W <= 256:\n        vol = np.transpose(vol, (2, 0, 1))\n    return vol\n\n\n# ============================================================\n# Helpers: norm / z-window / tiling\n# ============================================================\n\ndef norm_np(x: np.ndarray, mode: str) -> np.ndarray:\n    if mode == \"none\":\n        return x\n    if mode == \"zscore\":\n        m = float(x.mean())\n        s = float(x.std()) + 1e-6\n        return (x - m) / s\n    if mode == \"minmax\":\n        mn = float(x.min())\n        mx = float(x.max())\n        return (x - mn) / (mx - mn + 1e-6)\n    raise ValueError(f\"Unknown NORMALIZE={mode}\")\n\ndef _gather_z_window(vol: np.ndarray, d_center: int, z_window: int) -> np.ndarray:\n    D = vol.shape[0]\n    half = z_window // 2\n    idxs = []\n    for k in range(-half, half + 1):\n        dd = d_center + k\n        if dd < 0: dd = 0\n        if dd >= D: dd = D - 1\n        idxs.append(dd)\n    return vol[np.array(idxs, dtype=np.int64)]\n\ndef _autocast_ctx(enabled: bool):\n    if DEVICE.type != \"cuda\" or (not enabled):\n        return nullcontext()\n    dtype = torch.bfloat16 if USE_BF16 else torch.float16\n    try:\n        return torch.autocast(device_type=\"cuda\", dtype=dtype, enabled=True)\n    except TypeError:\n        from torch.cuda.amp import autocast as cuda_autocast\n        return cuda_autocast(dtype=dtype, enabled=True)\n\ndef make_weight_window(h: int, w: int, kind: str = \"hann\") -> np.ndarray:\n    if kind == \"ones\":\n        return np.ones((h, w), dtype=np.float32)\n    wy = np.hanning(h).astype(np.float32)\n    wx = np.hanning(w).astype(np.float32)\n    ww = np.outer(wy, wx)\n    return np.maximum(ww, 1e-3)\n\ndef _tile_grid(H: int, W: int, tile: int, stride: int) -> Tuple[List[int], List[int]]:\n    if stride <= 0:\n        stride = max(1, tile // 2)\n    ys = list(range(0, max(1, H - tile + 1), stride))\n    xs = list(range(0, max(1, W - tile + 1), stride))\n    if ys[-1] != H - tile:\n        ys.append(H - tile)\n    if xs[-1] != W - tile:\n        xs.append(W - tile)\n    return ys, xs\n\ndef _pad_hw_2d(arr2d: np.ndarray, crop: int, mode: str = \"edge\", constant: float = 0.0) -> Tuple[np.ndarray, Tuple[int,int,int,int]]:\n    H, W = arr2d.shape\n    pad_h = max(0, crop - H)\n    pad_w = max(0, crop - W)\n    pt = pad_h // 2\n    pb = pad_h - pt\n    pl = pad_w // 2\n    pr = pad_w - pl\n    if pad_h == 0 and pad_w == 0:\n        return arr2d, (0,0,0,0)\n    if mode == \"constant\":\n        out = np.pad(arr2d, ((pt,pb),(pl,pr)), mode=\"constant\", constant_values=constant)\n    else:\n        out = np.pad(arr2d, ((pt,pb),(pl,pr)), mode=mode)\n    return out, (pt,pb,pl,pr)\n\ndef _pad_hw_3d_zyx(vol: np.ndarray, crop: int, mode: str = \"edge\", constant: float = 0.0) -> Tuple[np.ndarray, Tuple[int,int,int,int]]:\n    # vol: (Z,H,W)\n    Z, H, W = vol.shape\n    pad_h = max(0, crop - H)\n    pad_w = max(0, crop - W)\n    pt = pad_h // 2\n    pb = pad_h - pt\n    pl = pad_w // 2\n    pr = pad_w - pl\n    if pad_h == 0 and pad_w == 0:\n        return vol, (0,0,0,0)\n    if mode == \"constant\":\n        out = np.pad(vol, ((0,0),(pt,pb),(pl,pr)), mode=\"constant\", constant_values=constant)\n    else:\n        out = np.pad(vol, ((0,0),(pt,pb),(pl,pr)), mode=mode)\n    return out, (pt,pb,pl,pr)\n\ndef _unpad_2d(arr2d: np.ndarray, pads: Tuple[int,int,int,int], orig_hw: Tuple[int,int]) -> np.ndarray:\n    pt,pb,pl,pr = pads\n    H, W = orig_hw\n    if pt==pb==pl==pr==0:\n        return arr2d\n    return arr2d[pt:pt+H, pl:pl+W]\n\n\n# ============================================================\n# Orientation transforms (canonical = ori1 = (Z,Y,X))\n# ============================================================\n\ndef orient_vol(vol: np.ndarray, orientation: int) -> np.ndarray:\n    o = int(orientation)\n    if o == 1:\n        return vol\n    if o == 2:\n        return np.transpose(vol, (1, 0, 2))  # (Z,Y,X)->(Y,Z,X)\n    if o == 3:\n        return np.transpose(vol, (2, 0, 1))  # (Z,Y,X)->(X,Z,Y)\n    raise ValueError(\"orientation must be 1,2,3\")\n\ndef unorient_vol(vol_oriented: np.ndarray, orientation: int, canonical_shape: Tuple[int,int,int]) -> np.ndarray:\n    o = int(orientation)\n    if o == 1:\n        out = vol_oriented\n    elif o == 2:\n        out = np.transpose(vol_oriented, (1, 0, 2))  # (Y,Z,X)->(Z,Y,X)\n    elif o == 3:\n        out = np.transpose(vol_oriented, (1, 2, 0))  # (X,Z,Y)->(Z,Y,X)\n    else:\n        raise ValueError(\"orientation must be 1,2,3\")\n\n    # safety crop\n    Z,Y,X = canonical_shape\n    out = out[:Z, :Y, :X]\n    return out\n\n\n# ============================================================\n# Model blocks (MATCH TRAINING de base HybridModel)\n# ============================================================\n\ndef is_power_of_two(x: int) -> bool:\n    return x > 0 and (x & (x - 1)) == 0\n\ndef _gn2d(ch: int) -> nn.GroupNorm:\n    for g in (16, 8, 4, 2):\n        if ch % g == 0:\n            return nn.GroupNorm(g, ch)\n    return nn.GroupNorm(1, ch)\n\ndef _gn3d(ch: int) -> nn.GroupNorm:\n    for g in (16, 8, 4, 2):\n        if ch % g == 0:\n            return nn.GroupNorm(g, ch)\n    return nn.GroupNorm(1, ch)\n\nclass DropPath(nn.Module):\n    def __init__(self, p: float = 0.0):\n        super().__init__()\n        self.p = float(p)\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        if self.p == 0.0 or (not self.training):\n            return x\n        keep = 1.0 - self.p\n        shape = (x.shape[0],) + (1,) * (x.ndim - 1)\n        rnd = torch.rand(shape, device=x.device, dtype=x.dtype)\n        mask = (rnd < keep).float()\n        return x * mask / keep\n\nclass MLP(nn.Module):\n    def __init__(self, dim: int, hidden_dim: Optional[int] = None, drop: float = 0.0):\n        super().__init__()\n        hidden_dim = hidden_dim or dim * 4\n        self.fc1 = nn.Linear(dim, hidden_dim)\n        self.act = nn.GELU()\n        self.fc2 = nn.Linear(hidden_dim, dim)\n        self.drop = nn.Dropout(drop)\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.drop(self.act(self.fc1(x)))\n        x = self.drop(self.fc2(x))\n        return x\n\nclass SelfAttentionSDPA(nn.Module):\n    def __init__(self, dim: int, heads: int, attn_drop: float = 0.0, proj_drop: float = 0.0):\n        super().__init__()\n        assert dim % heads == 0\n        self.heads = int(heads)\n        self.head_dim = dim // heads\n        self.qkv = nn.Linear(dim, dim * 3, bias=True)\n        self.attn_drop = float(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        B, N, C = x.shape\n        qkv = self.qkv(x).view(B, N, 3, self.heads, self.head_dim).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n        out = F.scaled_dot_product_attention(\n            q, k, v,\n            attn_mask=None,\n            dropout_p=self.attn_drop if self.training else 0.0,\n            is_causal=False,\n        )\n        out = out.transpose(1, 2).contiguous().view(B, N, C)\n        out = self.proj_drop(self.proj(out))\n        return out\n\nclass TransformerBlock(nn.Module):\n    def __init__(self, dim: int, heads: int, mlp_ratio: float = 4.0, drop: float = 0.0, drop_path: float = 0.0):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(dim)\n        self.attn = SelfAttentionSDPA(dim, heads, attn_drop=drop, proj_drop=drop)\n        self.drop_path = DropPath(drop_path)\n        self.norm2 = nn.LayerNorm(dim)\n        self.mlp = MLP(dim, hidden_dim=int(dim * mlp_ratio), drop=drop)\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = x + self.drop_path(self.attn(self.norm1(x)))\n        x = x + self.drop_path(self.mlp(self.norm2(x)))\n        return x\n\n# --- Residual + Attention blocks (scSE / SE) ---\nclass SEBlock2D(nn.Module):\n    def __init__(self, ch: int, reduction: int = 16):\n        super().__init__()\n        r = max(1, ch // int(reduction))\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(ch, r, 1, bias=True),\n            nn.ReLU(True),\n            nn.Conv2d(r, ch, 1, bias=True),\n            nn.Sigmoid(),\n        )\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        w = self.fc(self.pool(x))\n        return x * w\n\nclass SpatialSE2D(nn.Module):\n    def __init__(self, ch: int):\n        super().__init__()\n        self.conv = nn.Conv2d(ch, 1, kernel_size=1, bias=True)\n        self.act = nn.Sigmoid()\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        w = self.act(self.conv(x))\n        return x * w\n\nclass SCSEBlock2D(nn.Module):\n    def __init__(self, ch: int, reduction: int = 16):\n        super().__init__()\n        self.cSE = SEBlock2D(ch, reduction=reduction)\n        self.sSE = SpatialSE2D(ch)\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.cSE(x) + self.sSE(x)\n\nclass ResConvBlock2D(nn.Module):\n    def __init__(self, ic: int, oc: int, norm_fn, use_scse: bool = True, scse_reduction: int = 16):\n        super().__init__()\n        self.conv1 = nn.Conv2d(ic, oc, 3, padding=1, bias=False)\n        self.n1 = norm_fn(oc)\n        self.conv2 = nn.Conv2d(oc, oc, 3, padding=1, bias=False)\n        self.n2 = norm_fn(oc)\n        self.act = nn.ReLU(True)\n        self.skip = nn.Identity() if ic == oc else nn.Conv2d(ic, oc, 1, bias=False)\n        self.attn = SCSEBlock2D(oc, reduction=scse_reduction) if use_scse else nn.Identity()\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        identity = self.skip(x)\n        out = self.act(self.n1(self.conv1(x)))\n        out = self.n2(self.conv2(out))\n        out = self.attn(out)\n        out = self.act(out + identity)\n        return out\n\nclass SEBlock3D(nn.Module):\n    def __init__(self, ch: int, reduction: int = 16):\n        super().__init__()\n        r = max(1, ch // int(reduction))\n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Sequential(\n            nn.Conv3d(ch, r, 1, bias=True),\n            nn.ReLU(True),\n            nn.Conv3d(r, ch, 1, bias=True),\n            nn.Sigmoid(),\n        )\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        w = self.fc(self.pool(x))\n        return x * w\n\nclass ResConvBlock3D(nn.Module):\n    def __init__(self, ic: int, oc: int, norm_fn, use_se: bool = True, se_reduction: int = 16):\n        super().__init__()\n        self.conv1 = nn.Conv3d(ic, oc, 3, padding=1, bias=False)\n        self.n1 = norm_fn(oc)\n        self.conv2 = nn.Conv3d(oc, oc, 3, padding=1, bias=False)\n        self.n2 = norm_fn(oc)\n        self.act = nn.ReLU(True)\n        self.skip = nn.Identity() if ic == oc else nn.Conv3d(ic, oc, 1, bias=False)\n        self.attn = SEBlock3D(oc, reduction=se_reduction) if use_se else nn.Identity()\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        identity = self.skip(x)\n        out = self.act(self.n1(self.conv1(x)))\n        out = self.n2(self.conv2(out))\n        out = self.attn(out)\n        out = self.act(out + identity)\n        return out\n\ndef conv_block_2d(in_ch: int, out_ch: int) -> nn.Module:\n    return ResConvBlock2D(in_ch, out_ch, norm_fn=_gn2d, use_scse=True, scse_reduction=16)\n\nclass UpCatBlock2D(nn.Module):\n    def __init__(self, in_ch: int, skip_ch: int, out_ch: int):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)\n        self.conv = conv_block_2d(out_ch + skip_ch, out_ch)\n    def forward(self, x: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:\n        x = self.up(x)\n        if x.shape[-2:] != skip.shape[-2:]:\n            skip = F.interpolate(skip, size=x.shape[-2:], mode=\"bilinear\", align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        return self.conv(x)\n\nclass TransUNet25D_ViT(nn.Module):\n    \"\"\"\n    Input: (B,1,Z,H,W)\n    Output: (B,1,H,W)\n    \"\"\"\n    def __init__(\n        self,\n        z_window: int,\n        img_size: int,\n        patch_size: int,\n        overlap_patch: bool,\n        embed_dim: int,\n        depth: int,\n        heads: int,\n        drop: float,\n        drop_path: float,\n        base_channels: int,\n        refine_depth: int,\n    ):\n        super().__init__()\n        assert z_window % 2 == 1\n        assert is_power_of_two(patch_size)\n        self.z_window = int(z_window)\n        self.patch_size = int(patch_size)\n        self.overlap_patch = bool(overlap_patch)\n\n        self.stride = (patch_size // 2) if overlap_patch else patch_size\n        assert is_power_of_two(self.stride)\n        assert img_size % self.stride == 0, f\"img_size must be divisible by stride ({self.stride})\"\n\n        self.embed_dim = int(embed_dim)\n\n        stem_ch = int(base_channels)\n        self.stem3d = ResConvBlock3D(1, stem_ch, norm_fn=_gn3d, use_se=True, se_reduction=16)\n\n        self.collapse_conv = nn.Conv3d(stem_ch, stem_ch, kernel_size=(z_window, 1, 1), stride=1, padding=0, bias=False)\n        self.collapse_norm = _gn3d(stem_ch)\n        self.collapse_act  = nn.ReLU(inplace=True)\n        self.collapse_post = ResConvBlock3D(stem_ch, stem_ch, norm_fn=_gn3d, use_se=True, se_reduction=16)\n\n        self.num_ups = int(math.log2(self.stride))\n        chs = [base_channels * (2 ** i) for i in range(self.num_ups)]\n        bottleneck_ch = base_channels * (2 ** self.num_ups)\n\n        self.enc0 = conv_block_2d(stem_ch, chs[0])\n        self.downs = nn.ModuleList()\n        self.encs  = nn.ModuleList()\n        cur = chs[0]\n        for i in range(1, self.num_ups):\n            nxt = chs[i]\n            self.downs.append(nn.Conv2d(cur, nxt, 2, 2, bias=False))\n            self.encs.append(conv_block_2d(nxt, nxt))\n            cur = nxt\n\n        assert (patch_size - self.stride) % 2 == 0\n        patch_pad = (patch_size - self.stride) // 2\n        self.patch_embed = nn.Conv2d(stem_ch, embed_dim, kernel_size=patch_size, stride=self.stride, padding=patch_pad, bias=True)\n\n        g0 = img_size // self.stride\n        self.pos_embed = nn.Parameter(torch.zeros(1, embed_dim, g0, g0))\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n\n        dpr = torch.linspace(0, drop_path, steps=depth).tolist()\n        self.blocks = nn.ModuleList([TransformerBlock(embed_dim, heads, 4.0, drop, dpr[i]) for i in range(depth)])\n        self.norm = nn.LayerNorm(embed_dim)\n\n        self.to_bottleneck = nn.Sequential(\n            nn.Conv2d(embed_dim, bottleneck_ch, 1, bias=False),\n            _gn2d(bottleneck_ch),\n            nn.ReLU(inplace=True),\n        )\n\n        dec = []\n        cur = bottleneck_ch\n        for i in range(self.num_ups - 1, -1, -1):\n            out_ch = chs[i]\n            dec.append(UpCatBlock2D(cur, chs[i], out_ch))\n            cur = out_ch\n        self.decoder = nn.ModuleList(dec)\n\n        mods = []\n        for _ in range(int(refine_depth)):\n            mods += [nn.Conv2d(chs[0], chs[0], 3, padding=1, bias=False), _gn2d(chs[0]), nn.ReLU(inplace=True)]\n        self.refine = nn.Sequential(*mods)\n        self.head = nn.Conv2d(chs[0], 1, 1)\n\n    def _pos_tokens_2d(self, Hp: int, Wp: int) -> torch.Tensor:\n        pos = self.pos_embed\n        if pos.shape[-2:] != (Hp, Wp):\n            pos = F.interpolate(pos, size=(Hp, Wp), mode=\"bilinear\", align_corners=False)\n        return pos.flatten(2).transpose(1, 2)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        B, C, Z, H, W = x.shape\n        assert C == 1 and Z == self.z_window\n        assert H % self.stride == 0 and W % self.stride == 0\n\n        f3 = self.stem3d(x)\n        f2 = self.collapse_act(self.collapse_norm(self.collapse_conv(f3)))\n        f2 = self.collapse_post(f2)\n        f2 = f2.squeeze(2).contiguous()\n\n        s0 = self.enc0(f2)\n        skips = [s0]\n        cur = s0\n        for down, enc in zip(self.downs, self.encs):\n            cur = down(cur)\n            cur = enc(cur)\n            skips.append(cur)\n\n        feat = self.patch_embed(f2)\n        Hp, Wp = feat.shape[-2:]\n        tokens = feat.flatten(2).transpose(1, 2)\n        tokens = tokens + self._pos_tokens_2d(Hp, Wp)\n\n        for blk in self.blocks:\n            tokens = blk(tokens)\n        tokens = self.norm(tokens)\n\n        tmap = tokens.transpose(1, 2).reshape(B, self.embed_dim, Hp, Wp)\n        bott = self.to_bottleneck(tmap)\n\n        cur = bott\n        for i, upblk in enumerate(self.decoder):\n            skip = skips[-1 - i]\n            cur = upblk(cur, skip)\n\n        cur = self.refine(cur)\n        return self.head(cur)\n\ndef _gn(ch: int) -> nn.GroupNorm:\n    for g in (16, 8, 4, 2):\n        if ch % g == 0:\n            return nn.GroupNorm(g, ch)\n    return nn.GroupNorm(1, ch)\n\nclass AttentionGateGN(nn.Module):\n    def __init__(self, F_g: int, F_l: int, F_int: int):\n        super().__init__()\n        self.W_g = nn.Sequential(nn.Conv2d(F_g, F_int, 1, bias=False), _gn(F_int))\n        self.W_x = nn.Sequential(nn.Conv2d(F_l, F_int, 1, bias=False), _gn(F_int))\n        self.psi = nn.Sequential(nn.Conv2d(F_int, 1, 1, bias=False), _gn(1), nn.Sigmoid())\n        self.relu = nn.ReLU(True)\n    def forward(self, g, x):\n        psi = self.relu(self.W_g(g) + self.W_x(x))\n        psi = self.psi(psi)\n        return x * psi\n\nclass AttentionUNetGN(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, features=(64,128,256,512)):\n        super().__init__()\n        self.downs = nn.ModuleList()\n        self.pools = nn.ModuleList()\n        c = in_channels\n        for feat in features:\n            self.downs.append(ResConvBlock2D(c, feat, norm_fn=_gn, use_scse=True, scse_reduction=16))\n            self.pools.append(nn.MaxPool2d(2))\n            c = feat\n        self.bottleneck = ResConvBlock2D(features[-1], features[-1] * 2, norm_fn=_gn, use_scse=True, scse_reduction=16)\n        self.up_t = nn.ModuleList()\n        self.attn = nn.ModuleList()\n        self.up_c = nn.ModuleList()\n        rev = list(features)[::-1]\n        out = features[-1] * 2\n        for feat in rev:\n            self.up_t.append(nn.ConvTranspose2d(out, feat, 2, 2))\n            self.attn.append(AttentionGateGN(feat, feat, max(1, feat // 2)))\n            self.up_c.append(ResConvBlock2D(feat * 2, feat, norm_fn=_gn, use_scse=True, scse_reduction=16))\n            out = feat\n        self.final = nn.Conv2d(features[0], out_channels, 1)\n    def forward(self, x):\n        skips = []\n        out = x\n        for conv, pool in zip(self.downs, self.pools):\n            out = conv(out)\n            skips.append(out)\n            out = pool(out)\n        out = self.bottleneck(out)\n        skips = skips[::-1]\n        for up, att, conv in zip(self.up_t, self.attn, self.up_c):\n            out = up(out)\n            skip = att(out, skips.pop(0))\n            out = conv(torch.cat([out, skip], dim=1))\n        return self.final(out)\n\nclass GuidanceHead(nn.Module):\n    def __init__(self, in_ch: int = 1):\n        super().__init__()\n        self.b1 = ResConvBlock2D(in_ch, 32, norm_fn=_gn, use_scse=True, scse_reduction=16)\n        self.b2 = ResConvBlock2D(32, 16, norm_fn=_gn, use_scse=True, scse_reduction=16)\n        self.out = nn.Sequential(nn.Conv2d(16, 1, 1, bias=True), nn.Sigmoid())\n    def forward(self, x):\n        x = self.b1(x)\n        x = self.b2(x)\n        return self.out(x)\n\nclass HybridModel(nn.Module):\n    \"\"\"\n    MATCH TRAINING:\n      Stage2 in_ch=7\n      x2 = [z-1, z, z+1, prob1, attn, sobel(z), laplacian(z)]\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.stage1 = TransUNet25D_ViT(\n            z_window=Z_WINDOW,\n            img_size=CROP,\n            patch_size=PATCH,\n            overlap_patch=OVERLAP_PATCH,\n            embed_dim=EMBED,\n            depth=LAYERS,\n            heads=HEADS,\n            drop=DROP,\n            drop_path=DROPPATH,\n            base_channels=BASE_CH,\n            refine_depth=REFINE_DEPTH,\n        )\n        self.guidance = GuidanceHead(in_ch=1)\n        self.stage2 = AttentionUNetGN(in_channels=7, out_channels=1, features=REFINE_FEATURES)\n\n    def _edges2d_fp32(self, img: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        if img.device.type == \"cuda\":\n            with torch.amp.autocast(device_type=\"cuda\", enabled=False):\n                x = img.float()\n                kx = x.new_tensor([[-1,0,1],[-2,0,2],[-1,0,1]]).view(1,1,3,3)\n                ky = x.new_tensor([[-1,-2,-1],[0,0,0],[1,2,1]]).view(1,1,3,3)\n                lap = x.new_tensor([[0,1,0],[1,-4,1],[0,1,0]]).view(1,1,3,3)\n                gx = F.conv2d(x, kx, padding=1)\n                gy = F.conv2d(x, ky, padding=1)\n                sob = torch.sqrt(gx*gx + gy*gy + 1e-6)\n                l  = F.conv2d(x, lap, padding=1).abs()\n                return sob.to(dtype=img.dtype), l.to(dtype=img.dtype)\n        else:\n            x = img.float()\n            kx = x.new_tensor([[-1,0,1],[-2,0,2],[-1,0,1]]).view(1,1,3,3)\n            ky = x.new_tensor([[-1,-2,-1],[0,0,0],[1,2,1]]).view(1,1,3,3)\n            lap = x.new_tensor([[0,1,0],[1,-4,1],[0,1,0]]).view(1,1,3,3)\n            gx = F.conv2d(x, kx, padding=1)\n            gy = F.conv2d(x, ky, padding=1)\n            sob = torch.sqrt(gx*gx + gy*gy + 1e-6)\n            l  = F.conv2d(x, lap, padding=1).abs()\n            return sob.to(img.dtype), l.to(img.dtype)\n\n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        # x: (B,1,Z,H,W)\n        B, C, Z, H, W = x.shape\n        assert C == 1\n        zc = Z // 2\n        z_m1 = max(0, zc - 1)\n        z_p1 = min(Z - 1, zc + 1)\n\n        center0 = x[:, 0, zc].unsqueeze(1)    # (B,1,H,W)\n        c_m1    = x[:, 0, z_m1].unsqueeze(1)\n        c_p1    = x[:, 0, z_p1].unsqueeze(1)\n        center3 = torch.cat([c_m1, center0, c_p1], dim=1)  # (B,3,H,W)\n\n        logits1 = self.stage1(x)\n        prob1 = torch.sigmoid(logits1)\n        attn = self.guidance(prob1)\n\n        if DETACH_STAGE1_FOR_STAGE2:\n            prob1 = prob1.detach()\n            attn = attn.detach()\n\n        sob, lap = self._edges2d_fp32(center0)\n        x2 = torch.cat([center3, prob1, attn, sob, lap], dim=1)\n        logits2 = self.stage2(x2)\n        return logits1, logits2\n\n\n# ============================================================\n# Stacking head (MATCH TRAINER: StackingFusionUNetGN)\n# ============================================================\n\ndef _gn2d_big(ch: int) -> nn.GroupNorm:\n    for g in (32, 16, 8, 4, 2):\n        if ch % g == 0:\n            return nn.GroupNorm(g, ch)\n    return nn.GroupNorm(1, ch)\n\nclass ResBlockGN2(nn.Module):\n    def __init__(self, ic: int, oc: int, drop: float = 0.0):\n        super().__init__()\n        self.c1 = nn.Conv2d(ic, oc, 3, padding=1, bias=False)\n        self.n1 = _gn2d_big(oc)\n        self.c2 = nn.Conv2d(oc, oc, 3, padding=1, bias=False)\n        self.n2 = _gn2d_big(oc)\n        self.act = nn.SiLU(inplace=True)\n        self.skip = nn.Identity() if ic == oc else nn.Conv2d(ic, oc, 1, bias=False)\n        self.drop = nn.Dropout2d(drop) if drop and drop > 0 else nn.Identity()\n    def forward(self, x):\n        y = self.act(self.n1(self.c1(x)))\n        y = self.drop(y)\n        y = self.n2(self.c2(y))\n        return self.act(y + self.skip(x))\n\nclass Down2(nn.Module):\n    def __init__(self, ic: int, oc: int, drop: float = 0.0):\n        super().__init__()\n        self.pool = nn.AvgPool2d(2)\n        self.block = ResBlockGN2(ic, oc, drop=drop)\n    def forward(self, x):\n        return self.block(self.pool(x))\n\nclass Up2(nn.Module):\n    def __init__(self, ic: int, skipc: int, oc: int, drop: float = 0.0):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(ic, oc, 2, 2)\n        self.block = ResBlockGN2(oc + skipc, oc, drop=drop)\n    def forward(self, x, skip):\n        x = self.up(x)\n        if x.shape[-2:] != skip.shape[-2:]:\n            skip = F.interpolate(skip, size=x.shape[-2:], mode=\"bilinear\", align_corners=False)\n        return self.block(torch.cat([x, skip], dim=1))\n\nclass LogitCalibrator(nn.Module):\n    \"\"\"\n    Recalibra logits por orientación:\n      l' = scale*l + bias\n    \"\"\"\n    def __init__(self, n_models: int = 3, init_scale: float = 1.0):\n        super().__init__()\n        self.log_scale = nn.Parameter(torch.zeros(n_models))\n        self.bias = nn.Parameter(torch.zeros(n_models))\n        with torch.no_grad():\n            self.log_scale[:] = float(torch.log(torch.tensor(init_scale)))\n    def forward(self, logits3: torch.Tensor) -> torch.Tensor:\n        s = torch.exp(self.log_scale).view(1, -1, 1, 1)\n        b = self.bias.view(1, -1, 1, 1)\n        return logits3 * s + b\n\nclass StackingFusionUNetGN(nn.Module):\n    \"\"\"\n    forward(x25d, logits3) -> (out_logits, aux_dict)\n    x25d:   (B,Z,H,W)\n    logits3:(B,3,H,W)\n    \"\"\"\n    def __init__(\n        self,\n        z_window: int,\n        base: int = 64,\n        drop: float = 0.05,\n        use_boundary_head: bool = False,\n        use_disagreement: bool = True,\n        use_consensus_stats: bool = True,\n        use_probs: bool = True,\n        use_edges: bool = True,\n        use_x25d: bool = True,\n    ):\n        super().__init__()\n        assert z_window >= 3 and (z_window % 2 == 1)\n        self.z_window = int(z_window)\n        self.use_boundary = bool(use_boundary_head)\n\n        self.use_disagreement = bool(use_disagreement)\n        self.use_consensus_stats = bool(use_consensus_stats)\n        self.use_probs = bool(use_probs)\n        self.use_edges = bool(use_edges)\n        self.use_x25d = bool(use_x25d)\n\n        # edge kernels\n        kx = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=torch.float32).view(1,1,3,3)\n        ky = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]], dtype=torch.float32).view(1,1,3,3)\n        lap = torch.tensor([[0,1,0],[1,-4,1],[0,1,0]], dtype=torch.float32).view(1,1,3,3)\n        self.register_buffer(\"kx\", kx, persistent=False)\n        self.register_buffer(\"ky\", ky, persistent=False)\n        self.register_buffer(\"klap\", lap, persistent=False)\n\n        self.calib = LogitCalibrator(n_models=3, init_scale=1.0)\n\n        in_ch = 0\n        if self.use_x25d:\n            in_ch += self.z_window\n        in_ch += 3\n        if self.use_probs:\n            in_ch += 3\n        if self.use_consensus_stats:\n            in_ch += 4\n        if self.use_disagreement:\n            in_ch += 3\n        if self.use_edges:\n            in_ch += 2\n\n        c1 = base\n        c2 = base * 2\n        c3 = base * 4\n        c4 = base * 8\n\n        self.stem = ResBlockGN2(in_ch, c1, drop=drop)\n        self.d1 = Down2(c1, c2, drop=drop)\n        self.d2 = Down2(c2, c3, drop=drop)\n        self.d3 = Down2(c3, c4, drop=drop)\n\n        self.mid = nn.Sequential(\n            ResBlockGN2(c4, c4, drop=drop),\n            ResBlockGN2(c4, c4, drop=drop),\n        )\n\n        self.u3 = Up2(c4, c3, c3, drop=drop)\n        self.u2 = Up2(c3, c2, c2, drop=drop)\n        self.u1 = Up2(c2, c1, c1, drop=drop)\n\n        self.out = nn.Conv2d(c1, 1, 1)\n\n        if self.use_boundary:\n            self.out_bnd = nn.Conv2d(c1, 1, 1)\n\n    def _edges(self, center1: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        x = center1.float()\n        gx = F.conv2d(x, self.kx, padding=1)\n        gy = F.conv2d(x, self.ky, padding=1)\n        sob = torch.sqrt(gx*gx + gy*gy + 1e-6)\n        l = F.conv2d(x, self.klap, padding=1).abs()\n        return sob.to(center1.dtype), l.to(center1.dtype)\n\n    def forward(self, x25d: torch.Tensor, logits3: torch.Tensor):\n        B, Z, H, W = x25d.shape\n        if Z != self.z_window:\n            raise ValueError(f\"Expected Z_WINDOW={self.z_window}, got {Z}\")\n\n        l_raw = self.calib(logits3).clamp(-12.0, 12.0)\n        l_feat = torch.tanh(l_raw / 6.0)\n\n        feats = []\n        if self.use_x25d:\n            feats.append(x25d)\n        feats.append(l_feat)\n        if self.use_probs:\n            feats.append(torch.sigmoid(l_raw))\n        if self.use_consensus_stats:\n            mean = l_feat.mean(dim=1, keepdim=True)\n            var  = l_feat.var(dim=1, keepdim=True, unbiased=False)\n            mx   = l_feat.max(dim=1, keepdim=True).values\n            mn   = l_feat.min(dim=1, keepdim=True).values\n            feats += [mean, var, mx, mn]\n        if self.use_disagreement:\n            l1, l2, l3 = l_feat[:, 0:1], l_feat[:, 1:2], l_feat[:, 2:3]\n            feats += [(l1 - l2).abs(), (l1 - l3).abs(), (l2 - l3).abs()]\n        if self.use_edges:\n            zc = Z // 2\n            center1 = x25d[:, zc:zc+1]\n            sob, lap = self._edges(center1)\n            feats += [sob, lap]\n\n        x = torch.cat(feats, dim=1)\n\n        s1 = self.stem(x)\n        s2 = self.d1(s1)\n        s3 = self.d2(s2)\n        s4 = self.d3(s3)\n        m = self.mid(s4)\n        u3 = self.u3(m, s3)\n        u2 = self.u2(u3, s2)\n        u1 = self.u1(u2, s1)\n\n        out = self.out(u1)\n\n        aux = {\n            \"calib_scale\": torch.exp(self.calib.log_scale).detach(),\n            \"calib_bias\": self.calib.bias.detach(),\n        }\n        if self.use_boundary:\n            aux[\"bnd\"] = self.out_bnd(u1)\n        return out, aux\n\n\n# ============================================================\n# Checkpoint IO\n# ============================================================\n\ndef _torch_load(path: str, map_location=\"cpu\") -> Dict[str, Any]:\n    try:\n        return torch.load(path, map_location=map_location, weights_only=False)\n    except TypeError:\n        return torch.load(path, map_location=map_location)\n\ndef _load_state_dict_flexible(model: nn.Module, sd: Dict[str, Any]) -> None:\n    # maneja keys con/ sin \"module.\"\n    if any(k.startswith(\"module.\") for k in sd.keys()):\n        sd2 = {k.replace(\"module.\", \"\", 1): v for k, v in sd.items()}\n    else:\n        sd2 = sd\n    model.load_state_dict(sd2, strict=False)\n\ndef load_base_model(ckpt_path: str, device: torch.device) -> nn.Module:\n    assert Path(ckpt_path).exists(), f\"Missing checkpoint: {ckpt_path}\"\n    m = HybridModel().to(device)\n    ck = _torch_load(ckpt_path, map_location=\"cpu\")\n    sd = ck.get(\"model\", ck)\n    if isinstance(sd, dict) and \"state_dict\" in sd and isinstance(sd[\"state_dict\"], dict):\n        sd = sd[\"state_dict\"]\n    _load_state_dict_flexible(m, sd)\n    m.eval()\n    for p in m.parameters():\n        p.requires_grad_(False)\n    return m\n\ndef load_stack_model(ckpt_path: str, device: torch.device) -> nn.Module:\n    assert Path(ckpt_path).exists(), f\"Missing stacking checkpoint: {ckpt_path}\"\n    m = StackingFusionUNetGN(\n        z_window=Z_WINDOW,\n        base=int(FUSION_BASE),\n        drop=float(FUSION_DROP),\n        use_boundary_head=bool(FUSION_USE_BND),\n        use_disagreement=True,\n        use_consensus_stats=True,\n        use_probs=True,\n        use_edges=True,\n        use_x25d=True,\n    ).to(device)\n\n    ck = _torch_load(ckpt_path, map_location=\"cpu\")\n    sd = ck.get(\"head\", ck.get(\"model\", ck))\n    if isinstance(sd, dict) and \"state_dict\" in sd and isinstance(sd[\"state_dict\"], dict):\n        sd = sd[\"state_dict\"]\n\n    # prefer EMA shadow weights if present\n    if PREFER_EMA_WEIGHTS and isinstance(ck, dict) and isinstance(ck.get(\"ema\", None), dict):\n        ema_obj = ck[\"ema\"]\n        shadow = ema_obj.get(\"shadow\", None)\n        if isinstance(shadow, dict) and len(shadow) > 0:\n            # move to model dtype/device\n            cur = m.state_dict()\n            tmp = {}\n            for k, v in cur.items():\n                sv = shadow.get(k, None)\n                if sv is None:\n                    tmp[k] = v\n                else:\n                    tmp[k] = sv.to(device=device, dtype=v.dtype)\n            m.load_state_dict(tmp, strict=False)\n        else:\n            _load_state_dict_flexible(m, sd)\n    else:\n        _load_state_dict_flexible(m, sd)\n\n    m.eval()\n    for p in m.parameters():\n        p.requires_grad_(False)\n    return m\n\ndef resolve_threshold(stack_ckpt: str, ori1_ckpt: str) -> float:\n    if THRESHOLD >= 0:\n        return float(THRESHOLD)\n    # try stack ckpt\n    try:\n        ck = _torch_load(stack_ckpt, map_location=\"cpu\")\n        if isinstance(ck, dict):\n            thr = ck.get(\"val_best_thr\", None)\n            if thr is None and isinstance(ck.get(\"best_pack\", None), dict):\n                bp = ck[\"best_pack\"].get(\"best_by_best_dice\", None)\n                if isinstance(bp, dict) and (\"thr\" in bp):\n                    thr = bp[\"thr\"]\n            if thr is not None:\n                return float(thr)\n    except Exception:\n        pass\n    # try ori1 ckpt\n    try:\n        ck = _torch_load(ori1_ckpt, map_location=\"cpu\")\n        if isinstance(ck, dict):\n            thr = ck.get(\"val_best_thr\", ck.get(\"best_thr\", None))\n            if thr is not None:\n                return float(thr)\n    except Exception:\n        pass\n    return 0.5\n\n\n# ============================================================\n# Tiled inference: base logits2 volume\n# ============================================================\n\n@torch.no_grad()\ndef infer_logits2_slice_tiled_base(\n    model: nn.Module,\n    vol_oriented: np.ndarray,   # (D,H,W) float32\n    d_center: int,\n    device: torch.device,\n    tile: int,\n    stride: int,\n    ww: np.ndarray,\n) -> np.ndarray:\n    \"\"\"\n    Devuelve logits2 2D (H,W) para slice d_center sobre vol_oriented,\n    con tiles 2D + blending Hann.\n    \"\"\"\n    D, H, W = vol_oriented.shape\n\n    # pad a tile mínimo\n    vol2, pads = _pad_hw_3d_zyx(vol_oriented, tile, mode=\"edge\")\n    _, H2, W2 = vol2.shape\n\n    # precompute z-window once por slice\n    xz_full = _gather_z_window(vol2, d_center, Z_WINDOW)  # (Zwin,H2,W2)\n\n    acc = np.zeros((H2, W2), dtype=np.float32)\n    wsum = np.zeros((H2, W2), dtype=np.float32)\n\n    ys, xs = _tile_grid(H2, W2, tile, stride)\n\n    tiles_x = []\n    coords = []\n\n    use_amp = bool(USE_AMP and device.type == \"cuda\")\n\n    for y0 in ys:\n        for x0 in xs:\n            patch = xz_full[:, y0:y0+tile, x0:x0+tile]  # (Zwin,tile,tile)\n            patch = norm_np(patch.astype(np.float32, copy=False), NORMALIZE)\n            if CLAMP_AFTER_NORM is not None and float(CLAMP_AFTER_NORM) > 0:\n                patch = np.clip(patch, -float(CLAMP_AFTER_NORM), float(CLAMP_AFTER_NORM))\n\n            tiles_x.append(patch)\n            coords.append((y0, x0))\n\n            if len(tiles_x) >= int(BATCH_TILES):\n                xb = torch.from_numpy(np.stack(tiles_x, axis=0)).float().to(device)  # (B,Z,t,t)\n                xb = xb.unsqueeze(1)  # (B,1,Z,t,t)\n\n                with _autocast_ctx(use_amp):\n                    _, logits2 = model(xb)\n                out = logits2[:, 0].float().cpu().numpy()  # (B,t,t)\n\n                for i in range(out.shape[0]):\n                    yy, xx = coords[i]\n                    acc[yy:yy+tile, xx:xx+tile] += out[i] * ww\n                    wsum[yy:yy+tile, xx:xx+tile] += ww\n\n                tiles_x.clear()\n                coords.clear()\n\n    if tiles_x:\n        xb = torch.from_numpy(np.stack(tiles_x, axis=0)).float().to(device).unsqueeze(1)\n        with _autocast_ctx(use_amp):\n            _, logits2 = model(xb)\n        out = logits2[:, 0].float().cpu().numpy()\n        for i in range(out.shape[0]):\n            yy, xx = coords[i]\n            acc[yy:yy+tile, xx:xx+tile] += out[i] * ww\n            wsum[yy:yy+tile, xx:xx+tile] += ww\n\n    pred = acc / np.maximum(wsum, 1e-6)\n\n    # unpad a (H,W)\n    pred = _unpad_2d(pred, pads, (H, W))\n    return pred.astype(np.float32, copy=False)\n\ndef build_logits2_volume_canonical(\n    stem: str,\n    xvol_canon: np.ndarray,     # (Z,Y,X) float32\n    base_model: nn.Module,\n    orientation: int,\n    device: torch.device,\n    tmp_dir: Path,\n    tile: int,\n    stride: int,\n    weight_kind: str = \"hann\",\n) -> np.memmap:\n    \"\"\"\n    Construye logits2 volume canónico (Z,Y,X) para una orientación,\n    guardándolo en memmap float16 en disco. Devuelve memmap (lectura/escritura).\n    \"\"\"\n    tmp_dir.mkdir(parents=True, exist_ok=True)\n    Z, Y, X = xvol_canon.shape\n    out_path = tmp_dir / f\"{stem}_ori{orientation}_logits2_canon_f16.dat\"\n\n    # reuse si existe con mismo shape esperado\n    if out_path.exists():\n        mm = np.memmap(out_path, dtype=np.float16, mode=\"r\", shape=(Z, Y, X))\n        return mm\n\n    vol_o = orient_vol(xvol_canon, orientation).astype(np.float32, copy=False)\n    Do, Ho, Wo = vol_o.shape\n\n    # output oriented float16 memmap\n    out_o_path = tmp_dir / f\"{stem}_ori{orientation}_logits2_oriented_f16.dat\"\n    out_o = np.memmap(out_o_path, dtype=np.float16, mode=\"w+\", shape=(Do, Ho, Wo))\n\n    ww = make_weight_window(tile, tile, kind=weight_kind)\n\n    pbar = tqdm(range(Do), desc=f\"[BASE] ori{orientation} slices\", dynamic_ncols=True)\n    for d in pbar:\n        pred2d = infer_logits2_slice_tiled_base(\n            base_model, vol_o, d_center=d, device=device,\n            tile=tile, stride=stride, ww=ww\n        )\n        out_o[d] = pred2d.astype(np.float16, copy=False)\n\n    # unorient to canonical and save canonical memmap\n    out_c = np.memmap(out_path, dtype=np.float16, mode=\"w+\", shape=(Z, Y, X))\n\n    # streaming transpose to avoid big RAM\n    if orientation == 1:\n        out_c[:] = out_o[:]\n    elif orientation == 2:\n        # out_o: (Y,Z,X) -> out_c: (Z,Y,X)\n        for y in tqdm(range(out_o.shape[0]), desc=f\"[BASE] ori{orientation} unorient\", dynamic_ncols=True):\n            out_c[:, y, :] = out_o[y, :, :]\n    elif orientation == 3:\n        # out_o: (X,Z,Y) -> out_c: (Z,Y,X)\n        for x in tqdm(range(out_o.shape[0]), desc=f\"[BASE] ori{orientation} unorient\", dynamic_ncols=True):\n            out_c[:, :, x] = out_o[x, :, :]\n    else:\n        raise ValueError(\"orientation must be 1,2,3\")\n\n    # flush\n    out_c.flush()\n    try:\n        del out_o\n    except Exception:\n        pass\n    return np.memmap(out_path, dtype=np.float16, mode=\"r\", shape=(Z, Y, X))\n\n\n# ============================================================\n# Tiled inference: stacking slice\n# ============================================================\n\n@torch.no_grad()\ndef infer_stack_slice_tiled(\n    stack_model: nn.Module,\n    xvol_canon: np.ndarray,     # (Z,Y,X) float32\n    l1_canon: np.ndarray,       # (Z,Y,X) float16/float32\n    l2_canon: np.ndarray,\n    l3_canon: np.ndarray,\n    z: int,\n    device: torch.device,\n    tile: int,\n    stride: int,\n    ww: np.ndarray,\n) -> np.ndarray:\n    \"\"\"\n    Devuelve logits 2D (Y,X) para slice z (canónico),\n    usando tiles+Hann sobre stacking head:\n      x25d: (B,Zwin,tile,tile)\n      logits3: (B,3,tile,tile)\n    \"\"\"\n    Z, Y, X = xvol_canon.shape\n    z = int(z)\n\n    xz = _gather_z_window(xvol_canon, z, Z_WINDOW).astype(np.float32, copy=False)  # (Zwin,Y,X)\n    l1 = np.asarray(l1_canon[z], dtype=np.float32)\n    l2 = np.asarray(l2_canon[z], dtype=np.float32)\n    l3 = np.asarray(l3_canon[z], dtype=np.float32)\n\n    # pad to tile minimum (Y,X)\n    xz2, pads = _pad_hw_3d_zyx(xz, tile, mode=\"edge\")\n    l1p, _ = _pad_hw_2d(l1, tile, mode=\"edge\")\n    l2p, _ = _pad_hw_2d(l2, tile, mode=\"edge\")\n    l3p, _ = _pad_hw_2d(l3, tile, mode=\"edge\")\n\n    _, Y2, X2 = xz2.shape\n\n    acc = np.zeros((Y2, X2), dtype=np.float32)\n    wsum = np.zeros((Y2, X2), dtype=np.float32)\n\n    ys, xs = _tile_grid(Y2, X2, tile, stride)\n\n    bx = []\n    bl = []\n    coords = []\n\n    use_amp = bool(USE_AMP and device.type == \"cuda\")\n\n    for y0 in ys:\n        for x0 in xs:\n            patch_x = xz2[:, y0:y0+tile, x0:x0+tile]  # (Zwin,t,t)\n            patch_x = norm_np(patch_x, NORMALIZE)\n            if CLAMP_AFTER_NORM is not None and float(CLAMP_AFTER_NORM) > 0:\n                patch_x = np.clip(patch_x, -float(CLAMP_AFTER_NORM), float(CLAMP_AFTER_NORM))\n\n            patch_l1 = l1p[y0:y0+tile, x0:x0+tile]\n            patch_l2 = l2p[y0:y0+tile, x0:x0+tile]\n            patch_l3 = l3p[y0:y0+tile, x0:x0+tile]\n            patch_l = np.stack([patch_l1, patch_l2, patch_l3], axis=0).astype(np.float32, copy=False)  # (3,t,t)\n\n            bx.append(patch_x)\n            bl.append(patch_l)\n            coords.append((y0, x0))\n\n            if len(bx) >= int(BATCH_TILES):\n                x_b = torch.from_numpy(np.stack(bx, axis=0)).float().to(device)  # (B,Z,t,t)\n                l_b = torch.from_numpy(np.stack(bl, axis=0)).float().to(device)  # (B,3,t,t)\n\n                with _autocast_ctx(use_amp):\n                    out, _ = stack_model(x_b, l_b)\n                out_np = out[:, 0].float().cpu().numpy()  # (B,t,t)\n\n                for i in range(out_np.shape[0]):\n                    yy, xx = coords[i]\n                    acc[yy:yy+tile, xx:xx+tile] += out_np[i] * ww\n                    wsum[yy:yy+tile, xx:xx+tile] += ww\n\n                bx.clear()\n                bl.clear()\n                coords.clear()\n\n    if bx:\n        x_b = torch.from_numpy(np.stack(bx, axis=0)).float().to(device)\n        l_b = torch.from_numpy(np.stack(bl, axis=0)).float().to(device)\n        with _autocast_ctx(use_amp):\n            out, _ = stack_model(x_b, l_b)\n        out_np = out[:, 0].float().cpu().numpy()\n        for i in range(out_np.shape[0]):\n            yy, xx = coords[i]\n            acc[yy:yy+tile, xx:xx+tile] += out_np[i] * ww\n            wsum[yy:yy+tile, xx:xx+tile] += ww\n\n    pred = acc / np.maximum(wsum, 1e-6)\n    pred = _unpad_2d(pred, pads, (Y, X))\n    return pred.astype(np.float32, copy=False)\n\n\n# ============================================================\n# Outline helper (opcional)\n# ============================================================\n\ndef outline_from_mask(mask_u8: np.ndarray, conn8: bool = True) -> np.ndarray:\n    \"\"\"\n    mask_u8: (H,W) uint8 {0,1}\n    retorna outline uint8 {0,1}\n    \"\"\"\n    if mask_u8.dtype != np.uint8:\n        mask_u8 = mask_u8.astype(np.uint8, copy=False)\n    m = torch.from_numpy(mask_u8[None, None].astype(np.float32))\n    if conn8:\n        k = torch.ones((1,1,3,3), dtype=torch.float32)\n        need = 9.0\n    else:\n        k = torch.tensor([[[[0,1,0],[1,1,1],[0,1,0]]]], dtype=torch.float32)\n        need = 5.0\n    with torch.no_grad():\n        s = F.conv2d(m, k, padding=1)\n        er = (s >= need - 1e-3).float()\n        out = (m - er).clamp(min=0.0)\n    return out[0,0].byte().numpy()\n\n\n# ============================================================\n# Test volume listing\n# ============================================================\n\ndef list_test_items(test_dir: str) -> List[Tuple[str, str]]:\n    td = Path(test_dir)\n    if not td.exists():\n        raise FileNotFoundError(f\"TEST_DIR not found: {td}\")\n\n    tifs = sorted(td.glob(\"*.tif\")) + sorted(td.glob(\"*.tiff\"))\n    if tifs:\n        return [(p.stem, str(p)) for p in tifs]\n\n    items = []\n    for sub in sorted([p for p in td.iterdir() if p.is_dir()]):\n        # prefer common filename\n        cand = sub / \"surface_volume.tif\"\n        if cand.exists():\n            items.append((sub.name, str(cand)))\n            continue\n        # fallback: first tif\n        cands = sorted(sub.glob(\"*.tif\")) + sorted(sub.glob(\"*.tiff\"))\n        if cands:\n            items.append((sub.name, str(cands[0])))\n    return items\n\n\n# ============================================================\n# MAIN\n# ============================================================\n\ndef main():\n    out_dir = Path(OUT_DIR)\n    out_dir.mkdir(parents=True, exist_ok=True)\n\n    # stride default\n    stride = int(STRIDE) if int(STRIDE) > 0 else (CROP // 2)\n\n    # threshold resolve\n    thr = resolve_threshold(CKPT_STACK, CKPT_ORI1)\n    print(f\"[CFG] threshold={thr:.4f} (THRESHOLD={THRESHOLD}) | stride={stride} | amp={USE_AMP} bf16={USE_BF16}\")\n\n    # GPU setup\n    if DEVICE.type == \"cuda\":\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True\n        torch.backends.cudnn.benchmark = True\n        try:\n            torch.set_float32_matmul_precision(\"high\")\n        except Exception:\n            pass\n\n    # Load models\n    print(\"[LOAD] base models...\")\n    base1 = load_base_model(CKPT_ORI1, DEVICE)\n    base2 = load_base_model(CKPT_ORI2, DEVICE)\n    base3 = load_base_model(CKPT_ORI3, DEVICE)\n\n    print(\"[LOAD] stacking model...\")\n    stack = load_stack_model(CKPT_STACK, DEVICE)\n\n    # DataParallel\n    if USE_MULTI_GPU and DEVICE.type == \"cuda\" and torch.cuda.device_count() > 1:\n        ids = [i for i in GPU_IDS if i < torch.cuda.device_count()]\n        if len(ids) >= 2:\n            print(f\"[DP] enabling DataParallel on GPUs={ids}\")\n            base1 = nn.DataParallel(base1, device_ids=ids).to(DEVICE)\n            base2 = nn.DataParallel(base2, device_ids=ids).to(DEVICE)\n            base3 = nn.DataParallel(base3, device_ids=ids).to(DEVICE)\n            stack = nn.DataParallel(stack, device_ids=ids).to(DEVICE)\n\n    items = list_test_items(TEST_DIR)\n    if not items:\n        raise FileNotFoundError(f\"No test volumes found in {TEST_DIR}\")\n\n    print(f\"[TEST] found {len(items)} volumes\")\n\n    ww = make_weight_window(CROP, CROP, kind=\"hann\")\n    tmp_dir = out_dir / \"_tmp_predcache\"\n    tmp_dir.mkdir(parents=True, exist_ok=True)\n\n    saved_paths: List[Path] = []\n\n    for vid, vpath in items:\n        print(f\"\\n[VOL] {vid} -> {vpath}\")\n        vol = load_tiff_volume_no_codec(vpath)\n        vol = maybe_fix_dim_order(vol)\n        vol = vol.astype(np.float32, copy=False)\n\n        Z, Y, X = vol.shape\n        print(f\"[VOL] shape={vol.shape}\")\n\n        # Build three canonical logits2 volumes (memmap float16)\n        l1 = build_logits2_volume_canonical(vid, vol, base1, 1, DEVICE, tmp_dir, tile=CROP, stride=stride, weight_kind=\"hann\")\n        l2 = build_logits2_volume_canonical(vid, vol, base2, 2, DEVICE, tmp_dir, tile=CROP, stride=stride, weight_kind=\"hann\")\n        l3 = build_logits2_volume_canonical(vid, vol, base3, 3, DEVICE, tmp_dir, tile=CROP, stride=stride, weight_kind=\"hann\")\n\n        # Stacking inference per z\n        if OUTPUT_MODE == \"3d\":\n            out_mask = np.zeros((Z, Y, X), dtype=np.uint8)\n            z_iter = tqdm(range(Z), desc=f\"[STACK] {vid} z\", dynamic_ncols=True)\n            for z in z_iter:\n                logits2d = infer_stack_slice_tiled(\n                    stack, vol, l1, l2, l3, z=z,\n                    device=DEVICE, tile=CROP, stride=stride, ww=ww\n                )\n                probs = 1.0 / (1.0 + np.exp(-np.clip(logits2d, -20, 20)))\n                m = (probs > thr).astype(np.uint8)\n                if OUTLINE_ONLY:\n                    m = outline_from_mask(m, conn8=bool(OUTLINE_CONN8))\n                out_mask[z] = m\n        elif OUTPUT_MODE in (\"2d_mean\", \"2d_max\"):\n            # acumular probs\n            acc2d = None\n            z_iter = tqdm(range(Z), desc=f\"[STACK] {vid} z\", dynamic_ncols=True)\n            for z in z_iter:\n                logits2d = infer_stack_slice_tiled(\n                    stack, vol, l1, l2, l3, z=z,\n                    device=DEVICE, tile=CROP, stride=stride, ww=ww\n                )\n                probs = 1.0 / (1.0 + np.exp(-np.clip(logits2d, -20, 20)))\n                if acc2d is None:\n                    acc2d = probs\n                else:\n                    if OUTPUT_MODE == \"2d_mean\":\n                        acc2d += probs\n                    else:\n                        acc2d = np.maximum(acc2d, probs)\n            if OUTPUT_MODE == \"2d_mean\":\n                acc2d = acc2d / max(1, Z)\n            m2 = (acc2d > thr).astype(np.uint8)\n            if OUTLINE_ONLY:\n                m2 = outline_from_mask(m2, conn8=bool(OUTLINE_CONN8))\n            out_mask = m2  # 2D\n        else:\n            raise ValueError(f\"Unknown OUTPUT_MODE={OUTPUT_MODE}\")\n\n        # Save tif\n        out_path = out_dir / f\"{vid}.tif\"\n        if OUTPUT_MODE == \"3d\":\n            tifffile.imwrite(str(out_path), out_mask.astype(np.uint8), photometric=\"minisblack\")\n        else:\n            tifffile.imwrite(str(out_path), out_mask.astype(np.uint8), photometric=\"minisblack\")\n        saved_paths.append(out_path)\n        print(f\"[SAVE] {out_path}\")\n\n        if DEVICE.type == \"cuda\":\n            torch.cuda.empty_cache()\n\n    # Zip outputs\n    zip_path = Path(ZIP_PATH)\n    zip_path.parent.mkdir(parents=True, exist_ok=True)\n    with zipfile.ZipFile(str(zip_path), \"w\", compression=zipfile.ZIP_DEFLATED, compresslevel=6) as zf:\n        for p in saved_paths:\n            zf.write(str(p), arcname=p.name)\n    print(f\"\\n[ZIP] wrote {zip_path} with {len(saved_paths)} files\")\n    print(\"[DONE]\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-24T23:34:31.608624Z","iopub.execute_input":"2026-01-24T23:34:31.609252Z","iopub.status.idle":"2026-01-24T23:39:03.021266Z","shell.execute_reply.started":"2026-01-24T23:34:31.609213Z","shell.execute_reply":"2026-01-24T23:39:03.020404Z"}},"outputs":[],"execution_count":null}]}