{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":154281},{"sourceType":"datasetVersion","sourceId":18956429},{"sourceType":"datasetVersion","sourceId":18673646},{"sourceType":"datasetVersion","sourceId":18229736},{"sourceType":"datasetVersion","sourceId":18842180},{"sourceType":"datasetVersion","sourceId":18839182},{"sourceType":"datasetVersion","sourceId":18757740},{"sourceType":"datasetVersion","sourceId":18673450},{"sourceType":"datasetVersion","sourceId":18875869},{"sourceType":"datasetVersion","sourceId":18879001},{"sourceType":"datasetVersion","sourceId":18716507},{"sourceType":"kernelVersion","sourceId":342671664},{"sourceType":"kernelVersion","sourceId":342849430},{"sourceType":"modelInstanceVersion","sourceId":4533},{"sourceType":"datasetVersion","sourceId":18706996},{"sourceType":"datasetVersion","sourceId":18715672},{"sourceType":"datasetVersion","sourceId":19003959},{"sourceType":"modelInstanceVersion","sourceId":4534}],"dockerImageVersionId":31430,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"rsna_optimization":{"official_source_score":0.891,"revision":"v66-v65-parent-legacy-dino-002","source":"pilkwang/rsna-knee-baseline-v1"},"dinosaurs":{"name":"RSNA Knee | DINOsaur V5.5 TRAIN GoldOOFSlotHead","role":"training-only","output":"rsna-knee-dinov3-slothead-v16","next_step":"attach this notebook Output Files to V5.5 INFER"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"8ae3e66d-7a5c-4398-9470-249ea914f705","cell_type":"markdown","source":"# RSNA Knee | DINOsaur V4 — TRAIN 🦖\n\nExpected output folder: `/kaggle/working/rsna-knee-dinov3-slothead-v16/`.\nExpected files: `v16_manifest.json`, `v16_slothead_f*.pt`, `v16_slothead_oof.csv`.\n","metadata":{}},{"id":"3362f1f6-81ef-4ece-a28c-a8c00d6a7a84","cell_type":"code","source":"from pathlib import Path\nimport os, gc, json, math, time, hashlib, random, warnings, copy\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom sklearn.metrics import roc_auc_score\n\nwarnings.filterwarnings(\"ignore\")\nSEED = 20260820\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nOUT = Path(\"/kaggle/working/rsna-knee-dinov3-slothead-v16\")\nOUT.mkdir(parents=True, exist_ok=True)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"eb315541-8571-42b6-91ec-f061df4fd3ee","cell_type":"code","source":"import gc, os, time, warnings\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\nfrom pathlib import Path\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nwarnings.filterwarnings('ignore')\ncv2.setNumThreads(1)\nCROP_MM = 130.0\nSIZE = 336\nSLICE_BAND = (0.12, 0.88)\nN_SLICE = 16\nINTENSITY = 'slice'\nSLOTS = [('Sagittal', 1), ('Sagittal', 0), ('Coronal', 1), ('Coronal', 0), ('Axial', 1), ('Axial', 0)]\nN_SLOT = len(SLOTS)\nLABELS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\n\ndef _find_dir(*names):\n    root = Path('/kaggle/input')\n    cand = []\n    for n in names:\n        cand += [root / n, root / 'competitions' / n, root / 'datasets' / n]\n        for parent in (root / 'datasets', root / 'competitions', root):\n            if parent.is_dir():\n                try:\n                    cand += [d / n for d in parent.iterdir() if d.is_dir()]\n                except OSError:\n                    pass\n    for p in cand:\n        if p.is_dir():\n            return p\n    return None\nCOMP = _find_dir('rsna-knee-abnormality-detection')\nCKPT = _find_dir('knee-mri-fold-weights')\nassert COMP is not None, 'competition data not attached'\nassert CKPT is not None, 'fold weights not attached'\nassert (COMP / 'sample_submission.csv').exists(), f'no competition data at {COMP}'\nassert list(CKPT.glob('*_f*.pt')), f'no checkpoints at {CKPT}'\nDEV = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f'competition : {COMP}')\nprint(f'checkpoints : {CKPT}')\nprint(f'device      : {DEV}')\nfor i in range(torch.cuda.device_count() if DEV == 'cuda' else 0):\n    cc = torch.cuda.get_device_capability(i)\n    print(f'  gpu{i}       : {torch.cuda.get_device_name(i)} sm_{cc[0]}{cc[1]}, {torch.cuda.get_device_properties(i).total_memory / 2 ** 30:.0f} GiB, native bf16={cc >= (8, 0)}')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"6d1003a5-a11b-480a-943c-0c42ab2343d7","cell_type":"code","source":"SERIES_ROOT = COMP / 'test_series'\nif not SERIES_ROOT.exists():\n    SERIES_ROOT = COMP / 'train_series'\nprint('series root:', SERIES_ROOT)\n\ndef ordered_files(sdir, cap=64):\n    keyed = []\n    for f in sdir.glob('*.dcm'):\n        try:\n            ds = pydicom.dcmread(str(f), stop_before_pixels=True)\n            keyed.append((int(ds.InstanceNumber), str(f)))\n        except Exception:\n            continue\n        if len(keyed) >= cap * 4:\n            break\n    return [f for _, f in sorted(keyed)]\n\ndef series_side(path):\n    try:\n        return float(pydicom.dcmread(path, stop_before_pixels=True).ImagePositionPatient[0])\n    except Exception:\n        return 0.0\n\ndef read_crop(path):\n    try:\n        ds = pydicom.dcmread(path)\n        arr = ds.pixel_array.astype(np.float32)\n    except Exception:\n        return None\n    try:\n        ps = float(ds.PixelSpacing[0])\n    except Exception:\n        ps = CROP_MM / max(arr.shape)\n    half = int(round(CROP_MM / ps / 2))\n    cy, cx = (arr.shape[0] // 2, arr.shape[1] // 2)\n    y0, y1 = (max(0, cy - half), min(arr.shape[0], cy + half))\n    x0, x1 = (max(0, cx - half), min(arr.shape[1], cx + half))\n    crop = arr[y0:y1, x0:x1]\n    return None if crop.size == 0 else crop\n\ndef window(crop, lo, hi, flip):\n    c = np.clip((crop - lo) / max(hi - lo, 1e-06), 0, 1)\n    img = cv2.resize(c, (SIZE, SIZE), interpolation=cv2.INTER_AREA)\n    return img[:, ::-1].copy() if flip else img\n\ndef render(path, flip):\n    crop = read_crop(path)\n    if crop is None:\n        return None\n    lo, hi = np.percentile(crop[::4, ::4], [1, 99])\n    return window(crop, lo, hi, flip)\n\ndef build_study(args):\n    idx, study, recs = args\n    out = np.zeros((N_SLOT, N_SLICE, SIZE, SIZE), np.uint8)\n    mask = np.zeros(N_SLOT, np.uint8)\n    rows = pd.DataFrame(recs)\n    if len(rows):\n        for s_i, (plane, fs) in enumerate(SLOTS):\n            sub = rows[(rows.Anatomical_Plane == plane) & (rows.Fat_Suppression == fs)]\n            if sub.empty:\n                continue\n            files = ordered_files(SERIES_ROOT / study / sub.iloc[0].SeriesInstanceUID)\n            if not files:\n                continue\n            flip = plane != 'Sagittal' and series_side(files[0]) < 0\n            lo, hi = SLICE_BAND\n            i0 = int(round(lo * (len(files) - 1)))\n            i1 = int(round(hi * (len(files) - 1)))\n            avail = list(range(i0, i1 + 1))\n            if len(avail) >= N_SLICE:\n                picks = [avail[int(round(t))] for t in np.linspace(0, len(avail) - 1, N_SLICE)]\n                off = 0\n            else:\n                picks, off = (avail, (N_SLICE - len(avail)) // 2)\n            if INTENSITY == 'series':\n                crops = [read_crop(files[p]) for p in picks]\n                got = [x for x in crops if x is not None]\n                if got:\n                    samp = np.concatenate([x[::4, ::4].ravel() for x in got])\n                    lo_, hi_ = np.percentile(samp, [1, 99])\n                    for c, x in enumerate(crops):\n                        if x is None:\n                            x = read_crop(files[min(len(files) - 1, picks[c] + 1)])\n                        if x is not None:\n                            out[s_i, off + c] = (window(x, lo_, hi_, flip) * 255).astype(np.uint8)\n            else:\n                for c, p in enumerate(picks):\n                    img = render(files[p], flip)\n                    if img is None:\n                        img = render(files[min(len(files) - 1, p + 1)], flip)\n                    if img is not None:\n                        out[s_i, off + c] = (img * 255).astype(np.uint8)\n            mask[s_i] = len(picks)\n    return (idx, out, mask)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b7f330e9-a5be-4320-80a3-3c1ba3d2520a","cell_type":"code","source":"N_SLOT_TYPES, MASK_IDX = (6, 0)\n\ndef segment_softmax(scores, sidx, B):\n    T, K = scores.shape\n    idx = sidx.unsqueeze(1).expand(-1, K)\n    m = torch.full((B, K), float('-inf'), device=scores.device, dtype=scores.dtype)\n    m = m.scatter_reduce(0, idx, scores, reduce='amax', include_self=True)\n    e = (scores - m[sidx]).exp()\n    s = torch.zeros(B, K, device=scores.device, dtype=scores.dtype).index_add_(0, sidx, e)\n    return e / s[sidx].clamp(min=1e-06)\n\nclass MeanMaxPool(nn.Module):\n\n    def forward(self, f, sidx, B, slot=None, return_attn=False):\n        D = f.shape[1]\n        cnt = torch.zeros(B, device=f.device, dtype=f.dtype).index_add_(0, sidx, torch.ones(f.shape[0], device=f.device, dtype=f.dtype))\n        mean = torch.zeros(B, D, device=f.device, dtype=f.dtype).index_add_(0, sidx, f)\n        mean = mean / cnt.clamp(min=1).unsqueeze(1)\n        mx = torch.full((B, D), -10000.0, device=f.device, dtype=f.dtype)\n        mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), f, reduce='amax', include_self=True)\n        return (torch.cat([mean, mx], 1), None)\n\nclass LabelAttentionPool(nn.Module):\n\n    def __init__(self, d, n_labels=12, n_heads=4, slot_bias=True):\n        super().__init__()\n        self.d, self.k, self.h = (d, n_labels, n_heads)\n        self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n        self.key, self.val = (nn.Linear(d, d), nn.Linear(d, d))\n        self.slot_bias = nn.Parameter(torch.zeros(n_labels, N_SLOT_TYPES + 1)) if slot_bias else None\n\n    def forward(self, f, sidx, B, slot=None, return_attn=False):\n        scores = self.key(f) @ self.q.t() / self.d ** 0.5\n        if self.slot_bias is not None and slot is not None:\n            scores = scores + self.slot_bias.t()[slot]\n        a = segment_softmax(scores, sidx, B)\n        out = torch.zeros(B, self.k, self.d, device=f.device, dtype=f.dtype)\n        out = out.index_add_(0, sidx, a.unsqueeze(-1) * self.val(f).unsqueeze(1))\n        return (out, a)\n\nclass TokenXAttnPool(nn.Module):\n\n    def __init__(self, d, n_labels=12, n_heads=6, dropout=0.2):\n        super().__init__()\n        self.d, self.k = (d, n_labels)\n        self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n        self.slot_emb = nn.Embedding(N_SLOT_TYPES + 1, d, padding_idx=0)\n        self.kv_norm = nn.LayerNorm(d)\n        self.attn = nn.MultiheadAttention(d, n_heads, dropout=dropout, batch_first=True)\n\n    def forward(self, tok, sidx, B, slot=None, return_attn=False):\n        T, N, D = tok.shape\n        cnt = torch.bincount(sidx, minlength=B)\n        S = int(cnt.max().item())\n        starts = torch.cumsum(cnt, 0) - cnt\n        pos = torch.arange(T, device=tok.device) - starts[sidx]\n        kv = tok + self.slot_emb(slot).unsqueeze(1)\n        pad = tok.new_zeros(B, S, N, D)\n        pad[sidx, pos] = kv\n        keep = torch.zeros(B, S, dtype=torch.bool, device=tok.device)\n        keep[sidx, pos] = True\n        kpm = ~keep.repeat_interleave(N, dim=1)\n        pad = self.kv_norm(pad.reshape(B, S * N, D))\n        q = self.q.unsqueeze(0).expand(B, -1, -1)\n        att, w = self.attn(q, pad, pad, key_padding_mask=kpm, need_weights=return_attn, average_attn_weights=True)\n        cls = tok[:, 0]\n        mean = torch.zeros(B, D, device=tok.device, dtype=tok.dtype).index_add_(0, sidx, cls) / cnt.clamp(min=1).unsqueeze(1)\n        mx = torch.full((B, D), -10000.0, device=tok.device, dtype=tok.dtype)\n        mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), cls, reduce='amax', include_self=True)\n        base = torch.cat([mean, mx], 1).unsqueeze(1).expand(-1, self.k, -1)\n        return (torch.cat([att, base], -1), w)\n\nclass ViTSlotToken(nn.Module):\n\n    def __init__(self, vit, n_cat, dim=None):\n        super().__init__()\n        self.vit = vit\n        d = dim or vit.embed_dim\n        self.tok = nn.Embedding(n_cat + 1, d, padding_idx=MASK_IDX)\n        self.num_features = vit.num_features\n        self._orig_prefix = getattr(vit, 'num_prefix_tokens', 1)\n        vit.num_prefix_tokens = self._orig_prefix + 1\n        for blk in vit.blocks:\n            a = getattr(blk, 'attn', None)\n            if a is not None and hasattr(a, 'num_prefix_tokens'):\n                a.num_prefix_tokens = a.num_prefix_tokens + 1\n\n    @staticmethod\n    def _maybe(mod, x):\n        return x if mod is None else mod(x)\n\n    def forward_features(self, x, cat):\n        v = self.vit\n        x = v.patch_embed(x)\n        pos = v._pos_embed(x)\n        rope = None\n        if isinstance(pos, tuple):\n            x, rope = pos\n        else:\n            x = pos\n        x = self._maybe(getattr(v, 'patch_drop', None), x)\n        x = self._maybe(getattr(v, 'norm_pre', None), x)\n        npt = self._orig_prefix\n        tok = self.tok(cat).unsqueeze(1)\n        x = torch.cat([x[:, :npt], tok, x[:, npt:]], dim=1)\n        if rope is not None:\n            if getattr(v, 'rope_mixed', False):\n                for i, blk in enumerate(v.blocks):\n                    x = blk(x, rope=rope[i])\n            else:\n                for blk in v.blocks:\n                    x = blk(x, rope=rope)\n        else:\n            x = v.blocks(x)\n        return v.norm(x)\n\n    def forward_head(self, x, pre_logits=True):\n        return self.vit.forward_head(x, pre_logits=pre_logits)\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\nclass _GatedDepthBlock(nn.Module):\n\n    def __init__(self, n_slice, dropout=0.0, ls_init=0.1):\n        super().__init__()\n        self.norm = nn.GroupNorm(1, n_slice)\n        self.v = nn.Conv2d(n_slice, n_slice, 1)\n        self.g = nn.Conv2d(n_slice, n_slice, 1)\n        self.out = nn.Conv2d(n_slice, n_slice, 1)\n        self.gamma = nn.Parameter(torch.full((n_slice, 1, 1), ls_init))\n        self.drop = nn.Dropout2d(dropout) if dropout else nn.Identity()\n\n    def forward(self, x):\n        z = self.norm(x)\n        return x + self.gamma * self.drop(self.out(self.v(z) * F.silu(self.g(z))))\n\nclass DepthCompress(nn.Module):\n\n    def __init__(self, n_slice=16, out_ch=3, depth=1, dropout=0.0, ls_init=0.1, imagenet=True, proj_noise=0.25):\n        super().__init__()\n        self.imagenet = imagenet\n        self.blocks = nn.ModuleList([_GatedDepthBlock(n_slice, dropout, ls_init) for _ in range(depth)])\n        self.proj = nn.Conv2d(n_slice, out_ch, 1, bias=True)\n        if imagenet:\n            self.register_buffer('mu', torch.tensor(IMAGENET_MEAN).view(1, -1, 1, 1))\n            self.register_buffer('sd', torch.tensor(IMAGENET_STD).view(1, -1, 1, 1))\n\n    def forward(self, x):\n        keep = (x.amax(dim=1, keepdim=True) > 0).to(x.dtype)\n        z = x\n        for b in self.blocks:\n            z = b(z)\n        z = self.proj(z)\n        if self.imagenet:\n            z = (z - self.mu.to(z.dtype)) / self.sd.to(z.dtype)\n        return z * keep\nN_PLANE, N_CONTRAST = (3, 2)\n_PLANE_OF = lambda s: torch.clamp(s - 1, 0, 5) // 2\n_CONTRAST_OF = lambda s: torch.clamp(s - 1, 0, 5) % 2\n\nclass SlotDepthMixer(nn.Module):\n\n    def __init__(self, n_slice=16, ksize=5, alpha_max=0.25):\n        super().__init__()\n        self.n_slice, self.ksize, self.r = (n_slice, ksize, ksize // 2)\n        self.alpha_max = alpha_max\n        b = torch.tensor([1.0, 4.0, 6.0, 4.0, 1.0])\n        self.register_buffer('base', b.log()[self.r:])\n        n_u = self.r + 1\n        self.shared = nn.Parameter(torch.zeros(n_u))\n        self.plane_k = nn.Parameter(torch.zeros(N_PLANE, n_u))\n        self.contrast_k = nn.Parameter(torch.zeros(N_CONTRAST, n_u))\n        self.g0 = nn.Parameter(torch.zeros(()))\n        self.gate_p = nn.Parameter(torch.zeros(N_PLANE))\n        self.gate_c = nn.Parameter(torch.zeros(N_CONTRAST))\n        idx = torch.arange(n_slice)\n        self.register_buffer('off', idx[None, :] - idx[:, None])\n\n    def kernel(self, slot):\n        p, c = (_PLANE_OF(slot), _CONTRAST_OF(slot))\n        half = self.base + self.shared + self.plane_k[p] + self.contrast_k[c]\n        full = torch.cat([half.flip(-1)[..., :self.r], half], dim=-1)\n        return F.softmax(full, dim=-1)\n\n    def alpha(self, slot):\n        p, c = (_PLANE_OF(slot), _CONTRAST_OF(slot))\n        return self.alpha_max * torch.tanh(self.g0 + self.gate_p[p] + self.gate_c[c])\n\n    def forward(self, x, slot, vmask):\n        T, S, H, W = x.shape\n        if vmask is None:\n            raise ValueError('stem=mixer requires the padding mask')\n        k = self.kernel(slot)\n        v = vmask.to(k.dtype)\n        d = self.off + self.r\n        inb = (d >= 0) & (d < self.ksize)\n        kk = k[:, d.clamp(0, self.ksize - 1)] * inb\n        M = kk * v[:, None, :]\n        den = M.sum(-1, keepdim=True)\n        eye = torch.eye(S, device=x.device, dtype=M.dtype).expand(T, S, S)\n        ok = (den > 1e-06) & v[:, :, None].bool()\n        M = torch.where(ok, M / den.clamp(min=1e-06), eye)\n        a = self.alpha(slot)[:, None, None]\n        Aop = ((1.0 - a) * eye + a * M).to(x.dtype)\n        if x.is_contiguous(memory_format=torch.channels_last) and (not x.is_contiguous()):\n            y = torch.bmm(x.permute(0, 2, 3, 1).reshape(T, H * W, S), Aop.transpose(1, 2))\n            return y.reshape(T, H, W, S).permute(0, 3, 1, 2)\n        return torch.bmm(Aop, x.reshape(T, S, H * W)).reshape(T, S, H, W)\n\ndef _seg_mean_max(v, sidx, B):\n    D = v.shape[1]\n    cnt = torch.zeros(B, device=v.device, dtype=v.dtype).index_add_(0, sidx, torch.ones(v.shape[0], device=v.device, dtype=v.dtype))\n    mean = torch.zeros(B, D, device=v.device, dtype=v.dtype).index_add_(0, sidx, v)\n    mean = mean / cnt.clamp(min=1).unsqueeze(1)\n    mx = torch.full((B, D), -10000.0, device=v.device, dtype=v.dtype)\n    mx = mx.scatter_reduce(0, sidx.unsqueeze(1).expand(-1, D), v, reduce='amax', include_self=True)\n    return torch.cat([mean, mx], 1)\n\ndef _pad_kv(x, sidx, B, norm):\n    T, P, D = x.shape\n    cnt = torch.bincount(sidx, minlength=B)\n    S = int(cnt.max().item())\n    starts = torch.cumsum(cnt, 0) - cnt\n    pos = torch.arange(T, device=x.device) - starts[sidx]\n    pad = x.new_zeros(B, S, P, D)\n    pad[sidx, pos] = x\n    keep = torch.zeros(B, S, dtype=torch.bool, device=x.device)\n    keep[sidx, pos] = True\n    return (norm(pad.reshape(B, S * P, D)), ~keep.repeat_interleave(P, dim=1))\n\nclass _GatedDelta(nn.Module):\n\n    def __init__(self, d, n_labels, n_heads, dropout):\n        super().__init__()\n        self.q = nn.Parameter(torch.randn(n_labels, d) * 0.02)\n        self.kv_norm = nn.LayerNorm(d)\n        self.attn = nn.MultiheadAttention(d, n_heads, dropout=dropout, batch_first=True)\n        self.d_norm = nn.LayerNorm(d)\n        self.dw = nn.Parameter(torch.randn(n_labels, d) * (1.0 / d ** 0.5))\n        self.db = nn.Parameter(torch.zeros(n_labels))\n        self.gate = nn.Parameter(torch.zeros(n_labels))\n\n    def delta(self, pat, sidx, B, return_attn):\n        kv, kpm = _pad_kv(pat, sidx, B, self.kv_norm)\n        q = self.q.unsqueeze(0).expand(B, -1, -1)\n        att, w = self.attn(q, kv, kv, key_padding_mask=kpm, need_weights=return_attn, average_attn_weights=True)\n        return ((self.d_norm(att) * self.dw).sum(-1) + self.db, w)\n\nclass TokenResidualPool(_GatedDelta):\n\n    def __init__(self, d, n_labels=12, n_heads=6, pe=64, dropout=0.2):\n        super().__init__(d, n_labels, n_heads, dropout)\n        self.base = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(dropout), nn.Linear(2 * d + pe, n_labels))\n\n    def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n        base = self.base(torch.cat([_seg_mean_max(tok[:, 1:].mean(1), sidx, B), pres], 1))\n        d_, w = self.delta(tok[:, 1:], sidx, B, return_attn)\n        return (base + self.gate * d_, w)\n\nclass CodexResidualPool(_GatedDelta):\n\n    def __init__(self, d, n_labels=12, n_heads=6, pe=64, dropout=0.2):\n        super().__init__(d, n_labels, n_heads, dropout)\n        self.base = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(dropout), nn.Linear(2 * d + pe, n_labels))\n\n    def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n        base = self.base(torch.cat([_seg_mean_max(tok[:, 0], sidx, B), pres], 1))\n        d_, w = self.delta(tok[:, 1:], sidx, B, return_attn)\n        return (base + self.gate * d_, w)\n\nclass ClsAddPool(nn.Module):\n\n    def __init__(self, d, n_labels=12, pe=64, dropout=0.2):\n        super().__init__()\n        self.net = nn.Sequential(nn.LayerNorm(4 * d + pe), nn.Dropout(dropout), nn.Linear(4 * d + pe, n_labels))\n\n    def forward(self, tok, slot, sidx, B, pres, return_attn=False):\n        return (self.net(torch.cat([_seg_mean_max(tok[:, 1:].mean(1), sidx, B), _seg_mean_max(tok[:, 0], sidx, B), pres], 1)), None)\n\nclass Readout(nn.Module):\n\n    def __init__(self, pool, d, n_labels=12, pe=64):\n        super().__init__()\n        self.pool_kind, self.k = (pool, n_labels)\n        self.pres_emb = nn.Embedding(N_SLOT_TYPES + 1, pe, padding_idx=0)\n        if pool in ('xres', 'clsadd', 'xcodex'):\n            self.pool = {'xres': TokenResidualPool, 'clsadd': ClsAddPool, 'xcodex': CodexResidualPool}[pool](d, n_labels, pe=pe)\n        elif pool in ('attn', 'xattn'):\n            if pool == 'xattn':\n                self.pool = TokenXAttnPool(d, n_labels)\n                wd = 3 * d + pe\n            else:\n                self.pool = LabelAttentionPool(d, n_labels)\n                wd = d + pe\n            self.norm = nn.LayerNorm(wd)\n            self.w = nn.Parameter(torch.randn(n_labels, wd) * (1.0 / wd ** 0.5))\n            self.b = nn.Parameter(torch.zeros(n_labels))\n        else:\n            self.pool = MeanMaxPool()\n            self.net = nn.Sequential(nn.LayerNorm(2 * d + pe), nn.Dropout(0.2), nn.Linear(2 * d + pe, n_labels))\n        self.drop = nn.Dropout(0.2)\n\n    def forward(self, f, slot, sidx, B, return_attn=False):\n        pe = self.pres_emb(slot)\n        pres = torch.zeros(B, pe.shape[1], device=f.device, dtype=f.dtype).index_add_(0, sidx, pe)\n        if self.pool_kind in ('xres', 'clsadd', 'xcodex'):\n            return self.pool(f, slot, sidx, B, pres)[0]\n        pooled, attn = self.pool(f, sidx, B, slot=slot, return_attn=return_attn)\n        if self.pool_kind in ('attn', 'xattn'):\n            x = torch.cat([pooled, pres.unsqueeze(1).expand(-1, self.k, -1)], -1)\n            x = self.drop(self.norm(x))\n            return (x * self.w).sum(-1) + self.b\n        return self.net(torch.cat([pooled, pres], 1))\n\nclass Net(nn.Module):\n\n    def __init__(self, enc, cond, n_meta=0, pool='mean_max', stem='native', n_slice=16):\n        super().__init__()\n        self.enc, self.cond = (enc, cond)\n        self.compress = DepthCompress(n_slice, 3) if stem == 'compress' else None\n        self.mixer = SlotDepthMixer(n_slice) if stem == 'mixer' else None\n        self.tokens = pool in ('xattn', 'xres', 'clsadd', 'xcodex')\n        D = enc.num_features\n        self.meta_mlp = nn.Sequential(nn.LayerNorm(n_meta), nn.Linear(n_meta, 128), nn.GELU(), nn.Linear(128, D)) if n_meta > 0 else None\n        self.readout = Readout(pool, D)\n        if cond == 'post':\n            self.slot_emb = nn.Embedding(N_SLOT_TYPES + 1, D, padding_idx=MASK_IDX)\n\n    def forward(self, im, slot, smeta, sidx, B, vm=None):\n        if self.mixer is not None:\n            im = self.mixer(im, slot, vm)\n        if self.compress is not None:\n            im = self.compress(im)\n        f = self.enc.forward_features(im, slot) if self.cond == 'token' else self.enc.forward_features(im)\n        if self.tokens:\n            inner = getattr(self.enc, 'vit', self.enc)\n            orig = getattr(self.enc, '_orig_prefix', getattr(inner, 'num_prefix_tokens', 1))\n            f = torch.cat([f[:, :1], f[:, orig:]], 1)\n        else:\n            f = self.enc.forward_head(f, pre_logits=True)\n            if f.dim() > 2:\n                f = f.flatten(1)\n        ex = (lambda v: v.unsqueeze(1)) if self.tokens else lambda v: v\n        if self.cond == 'post':\n            f = f + ex(self.slot_emb(slot))\n        if self.meta_mlp is not None and smeta.shape[1] > 0:\n            mt = self.meta_mlp(smeta)\n            f = torch.cat([f, mt.unsqueeze(1)], 1) if self.tokens else f + mt\n        return self.readout(f, slot, sidx, B)\nmodels = []\nfor ckpt_path in sorted(CKPT.glob('*_f*.pt')):\n    z = torch.load(ckpt_path, map_location='cpu', weights_only=False)\n    cfg = z['cfg']\n    _stem = cfg.get('stem', 'native')\n    _in = 3 if _stem == 'compress' else cfg.get('n_slice', 16)\n    enc = timm.create_model(cfg['backbone'], pretrained=False, num_classes=0, in_chans=_in, **{'img_size': cfg['img']} if 'vit_' in cfg['backbone'] else {})\n    if cfg['cond'] == 'token':\n        enc = ViTSlotToken(enc, N_SLOT_TYPES)\n    m = Net(enc, cfg['cond'], cfg.get('n_meta', 0), cfg['pool'], stem=_stem, n_slice=cfg.get('n_slice', 16))\n    missing, unexpected = m.load_state_dict(z['state_dict'], strict=False)\n    assert not [k for k in missing if not k.startswith('enc.')], f'missing {missing[:5]}'\n    assert not unexpected, f'unexpected {unexpected[:5]}'\n    models.append(m.eval())\n    print(f\"loaded {ckpt_path.name}  fold {z['fold']}  {cfg['backbone']} pool={cfg['pool']} meta={cfg['meta']}\")\nCFG = cfg\nassert CFG.get('n_meta', 0) == 0, f\"checkpoint expects {CFG['n_meta']} metadata features -- build slot_meta for the TEST studies and pass it to predict() before submitting\"\nprint(f\"\\n{len(models)} fold models ready | input norm: {CFG.get('norm', 'none')}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"e3ab6d47-c1cc-40fd-9f03-495f926af937","cell_type":"code","source":"# -------------------- V16 frozen DINOv3 feature cache --------------------\n\nSERIES_ROOT = COMP / \"train_series\"\ntrain_df = pd.read_csv(\n    COMP / \"train.csv\",\n    dtype={\"StudyInstanceUID\": str},\n)\ntrain_series = pd.read_csv(\n    COMP / \"train_series.csv\",\n    dtype={\n        \"StudyInstanceUID\": str,\n        \"SeriesInstanceUID\": str,\n    },\n)\ntrain_series = train_series.loc[:, ~train_series.columns.duplicated()]\ntrain_ids = train_df[\"StudyInstanceUID\"].astype(str).tolist()\ntrain_by = {\n    str(uid): group.to_dict(\"records\")\n    for uid, group in train_series[\n        train_series[\"StudyInstanceUID\"].astype(str).isin(set(train_ids))\n    ].groupby(\"StudyInstanceUID\")\n}\n\ndef _v16_find_file(filename):\n    roots = [Path(\"/kaggle/input\"), Path(\"/kaggle/working\")]\n    matches = []\n    for root in roots:\n        if not root.exists():\n            continue\n        for current, dirs, files in os.walk(root):\n            dirs[:] = [\n                d for d in dirs\n                if d not in (\"train_series\", \"test_series\", \".git\", \"__pycache__\")\n            ]\n            if filename in files:\n                matches.append(Path(current) / filename)\n    matches = [p for p in matches if p.is_file()]\n    matches.sort(key=lambda p: (len(p.parts), str(p)))\n    return matches[0] if matches else None\n\ndef _v16_load_supervision():\n    oof_path = _v16_find_file(\"oof.npz\")\n    fold_path = _v16_find_file(\"v52_e11_oof.csv\")\n    public_path = _v16_find_file(\"v52_oof.csv\")\n\n    if oof_path is None or fold_path is None or public_path is None:\n        raise FileNotFoundError(\n            \"V16 needs oof.npz + v52_oof.csv + v52_e11_oof.csv as Kaggle inputs\"\n        )\n\n    with np.load(oof_path, allow_pickle=False) as data:\n        ids = data[\"ids\"].astype(str)\n        targets = data[\"targets\"].astype(str).tolist()\n        weak_y = data[\"y_derived\"].astype(np.float32)\n        base_oof = data[\"pred\"].astype(np.float32)\n        gold = data[\"gold_mask\"].astype(bool)\n\n    if targets != LABELS:\n        raise RuntimeError(\"OOF target contract mismatch\")\n    if not np.array_equal(ids, np.asarray(train_ids, dtype=str)):\n        raise RuntimeError(\"OOF train ID order mismatch\")\n\n    exact = train_df[LABELS].apply(pd.to_numeric, errors=\"coerce\")\n    official_gold = exact.notna().all(axis=1).to_numpy()\n    if not np.array_equal(gold, official_gold):\n        raise RuntimeError(\"gold_mask mismatch\")\n\n    y = np.where(np.isfinite(weak_y), weak_y, 0.5).astype(np.float32)\n    y[gold] = exact.loc[gold, LABELS].to_numpy(np.float32)\n\n    fold_frame = pd.read_csv(\n        fold_path,\n        dtype={\"StudyInstanceUID\": str},\n    )\n    fold_frame = train_df[[\"StudyInstanceUID\"]].merge(\n        fold_frame[[\"StudyInstanceUID\", \"fold\"] + LABELS],\n        on=\"StudyInstanceUID\",\n        how=\"left\",\n        validate=\"one_to_one\",\n    )\n    fold = pd.to_numeric(fold_frame[\"fold\"], errors=\"coerce\").to_numpy()\n    if not np.isfinite(fold).all():\n        raise RuntimeError(\"OOF fold assignment incomplete\")\n    fold = fold.astype(np.int64)\n\n    public_frame = pd.read_csv(\n        public_path,\n        dtype={\"StudyInstanceUID\": str},\n    )\n    public_frame = train_df[[\"StudyInstanceUID\"]].merge(\n        public_frame[[\"StudyInstanceUID\"] + LABELS],\n        on=\"StudyInstanceUID\",\n        how=\"left\",\n        validate=\"one_to_one\",\n    )\n\n    def rank_cols(a):\n        return pd.DataFrame(a).rank(method=\"average\", pct=True).to_numpy(np.float64)\n\n    base_rank = rank_cols(base_oof)\n    public_rank = rank_cols(public_frame[LABELS].to_numpy(np.float64))\n    pass2_rank = rank_cols(fold_frame[LABELS].to_numpy(np.float64))\n    approx_anchor = rank_cols(\n        0.85 * rank_cols(0.50 * base_rank + 0.50 * public_rank)\n        + 0.15 * pass2_rank\n    )\n\n    return y, gold, fold, approx_anchor, str(oof_path)\n\nY, GOLD, FOLD, APPROX_ANCHOR, OOF_SOURCE = _v16_load_supervision()\n\n# The encoder is exactly the already attached DINOv3 fold-0 model.\nencoder_model = models[0].eval()\nencoder_cfg = dict(CFG)\nencoder_ckpt = sorted(CKPT.glob(\"*_f*.pt\"))[0]\n\nDEV_FEAT = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nencoder_model = encoder_model.to(DEV_FEAT)\nAMP_DTYPE = torch.float16 if str(DEV_FEAT).startswith(\"cuda\") else torch.float32\n\ndef _v16_norm(im):\n    kind = encoder_cfg.get(\"norm\", \"none\")\n    if kind == \"zscore\":\n        m = (im > 0).float()\n        n = m.sum(dim=(1, 2, 3), keepdim=True).clamp(min=1.0)\n        mu = (im * m).sum(dim=(1, 2, 3), keepdim=True) / n\n        var = (((im - mu) * m) ** 2).sum(dim=(1, 2, 3), keepdim=True) / n\n        return (im - mu) / (var.sqrt() + 1e-6) * m\n    if kind == \"imagenet\":\n        m = (im > 0).float()\n        return (im - 0.485) / 0.229 * m\n    return im\n\n@torch.no_grad()\ndef _v16_encode_slots(images, masks, micro=8):\n    B = len(masks)\n    records = []\n\n    for study_index in range(B):\n        present = np.nonzero(masks[study_index] > 0)[0]\n        for slot in present:\n            records.append((study_index, int(slot)))\n\n    if not records:\n        return None, np.zeros((B, N_SLOT), np.uint8)\n\n    feature_batches = []\n    index_batches = []\n\n    for start in range(0, len(records), micro):\n        part = records[start:start + micro]\n        ims = np.stack(\n            [images[study_index, slot] for study_index, slot in part],\n            axis=0,\n        )\n        im = torch.from_numpy(ims).to(DEV_FEAT).float().div_(255.0)\n        im = _v16_norm(im)\n\n        slots = torch.tensor(\n            [slot + 1 for _, slot in part],\n            device=DEV_FEAT,\n            dtype=torch.long,\n        )\n        vm = torch.from_numpy(\n            ims.reshape(len(ims), ims.shape[1], -1).max(2) > 0\n        ).to(DEV_FEAT)\n\n        model = encoder_model\n        x = im\n        if model.mixer is not None:\n            x = model.mixer(x, slots, vm)\n        if model.compress is not None:\n            x = model.compress(x)\n\n        with torch.autocast(\n            \"cuda\" if str(DEV_FEAT).startswith(\"cuda\") else \"cpu\",\n            dtype=AMP_DTYPE,\n            enabled=str(DEV_FEAT).startswith(\"cuda\"),\n        ):\n            if encoder_cfg[\"cond\"] == \"token\":\n                tokens = model.enc.forward_features(x, slots)\n            else:\n                tokens = model.enc.forward_features(x)\n\n        if tokens.ndim == 2:\n            cls = tokens.float()\n            mean = tokens.float()\n            focal = tokens.float()\n        else:\n            inner = getattr(model.enc, \"vit\", model.enc)\n            prefix = getattr(\n                model.enc,\n                \"_orig_prefix\",\n                getattr(inner, \"num_prefix_tokens\", 1),\n            )\n            cls = tokens[:, 0].float()\n            patches = tokens[:, prefix:].float()\n            mean = patches.mean(dim=1)\n\n            n_patch = patches.shape[1]\n            grid = int(round(math.sqrt(n_patch)))\n            if grid * grid == n_patch and grid >= 3:\n                p = patches.reshape(len(part), grid, grid, -1)\n                radius = max(1, grid // 4)\n                center = grid // 2\n                y0, y1 = max(0, center - radius), min(grid, center + radius + 1)\n                focal = p[:, y0:y1, y0:y1].mean(dim=(1, 2))\n            else:\n                focal = mean\n\n        feat = torch.cat([cls, mean, focal], dim=1).cpu().numpy().astype(np.float16)\n        feature_batches.append(feat)\n        index_batches.extend(part)\n\n    feature_dim = feature_batches[0].shape[1]\n    out = np.zeros((B, N_SLOT, feature_dim), np.float16)\n    out_mask = np.zeros((B, N_SLOT), np.uint8)\n\n    offset = 0\n    for batch in feature_batches:\n        for row in range(len(batch)):\n            study_index, slot = index_batches[offset]\n            out[study_index, slot] = batch[row]\n            out_mask[study_index, slot] = 1\n            offset += 1\n\n    return out, out_mask\n\nFEATURE_CACHE = OUT / \"train_features_fp16.npz\"\n\nif FEATURE_CACHE.exists():\n    cache = np.load(FEATURE_CACHE, allow_pickle=False)\n    FEATURES = cache[\"features\"]\n    SLOT_MASK = cache[\"mask\"]\n    cached_ids = cache[\"ids\"].astype(str)\n    if not np.array_equal(cached_ids, np.asarray(train_ids, dtype=str)):\n        raise RuntimeError(\"feature cache ID mismatch\")\n    print(\"loaded feature cache\", FEATURES.shape)\nelse:\n    WORKERS = max(1, min(6, os.cpu_count() or 4))\n    CHUNK = 24\n    chunks = []\n    mask_chunks = []\n\n    t0 = time.time()\n    with ProcessPoolExecutor(max_workers=WORKERS) as ex:\n        for c0 in range(0, len(train_ids), CHUNK):\n            block = train_ids[c0:c0 + CHUNK]\n            imgs = np.zeros((len(block), N_SLOT, N_SLICE, SIZE, SIZE), np.uint8)\n            msks = np.zeros((len(block), N_SLOT), np.uint8)\n\n            futures = [\n                ex.submit(build_study, (i, uid, train_by.get(uid, [])))\n                for i, uid in enumerate(block)\n            ]\n            for future in as_completed(futures):\n                try:\n                    i, arr, mask = future.result()\n                    imgs[i], msks[i] = arr, mask\n                except Exception as exc:\n                    print(\"decode failed\", type(exc).__name__, exc)\n\n            feat, feat_mask = _v16_encode_slots(imgs, msks, micro=8)\n            if feat is None:\n                raise RuntimeError(\"empty feature block\")\n\n            chunks.append(feat)\n            mask_chunks.append(feat_mask)\n\n            done = min(c0 + len(block), len(train_ids))\n            elapsed = time.time() - t0\n            print(\n                f\"features {done}/{len(train_ids)} \"\n                f\"{elapsed/60:.1f}m \"\n                f\"eta={elapsed/max(done,1)*(len(train_ids)-done)/60:.1f}m\",\n                flush=True,\n            )\n\n            del imgs, msks, feat, feat_mask\n            gc.collect()\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n\n    FEATURES = np.concatenate(chunks, axis=0)\n    SLOT_MASK = np.concatenate(mask_chunks, axis=0)\n\n    np.savez_compressed(\n        FEATURE_CACHE,\n        ids=np.asarray(train_ids, dtype=\"U80\"),\n        features=FEATURES,\n        mask=SLOT_MASK,\n    )\n    print(\"saved feature cache\", FEATURES.shape)\n\nencoder_model = encoder_model.cpu()\ndel models\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c965df0f-a5bd-4b28-8928-7d4f63bb287f","cell_type":"code","source":"# -------------------- V16 target-conditioned slot-attention head --------------------\n\nclass SlotAttentionHead(nn.Module):\n    def __init__(self, feature_dim, hidden=384, n_target=12, n_slot=6):\n        super().__init__()\n        self.feature_dim = int(feature_dim)\n        self.hidden = int(hidden)\n        self.n_target = int(n_target)\n        self.n_slot = int(n_slot)\n\n        self.proj = nn.Sequential(\n            nn.LayerNorm(feature_dim),\n            nn.Linear(feature_dim, hidden),\n            nn.GELU(),\n        )\n        self.slot_emb = nn.Parameter(torch.zeros(n_slot, hidden))\n        nn.init.normal_(self.slot_emb, std=0.02)\n\n        self.query = nn.Parameter(torch.zeros(n_target, hidden))\n        nn.init.normal_(self.query, std=0.02)\n\n        self.attn = nn.MultiheadAttention(\n            hidden,\n            num_heads=6,\n            dropout=0.10,\n            batch_first=True,\n        )\n\n        self.fuse = nn.Sequential(\n            nn.LayerNorm(hidden * 5),\n            nn.Linear(hidden * 5, hidden),\n            nn.GELU(),\n            nn.Dropout(0.16),\n        )\n\n        self.target_weight = nn.Parameter(torch.empty(n_target, hidden))\n        self.target_bias = nn.Parameter(torch.zeros(n_target))\n        nn.init.normal_(self.target_weight, std=0.03)\n\n    def forward(\n        self,\n        features,\n        slot_mask,\n        modality_dropout=0.0,\n        feature_noise=0.0,\n    ):\n        mask = slot_mask.bool()\n\n        if self.training and modality_dropout > 0:\n            keep = torch.rand_like(slot_mask.float()) > float(modality_dropout)\n            keep = keep & mask\n            empty = keep.sum(dim=1) == 0\n            if empty.any():\n                first = mask[empty].float().argmax(dim=1)\n                keep[empty] = False\n                keep[empty, first] = True\n            mask = keep\n\n        x = self.proj(features.float())\n        x = x + self.slot_emb.unsqueeze(0)\n\n        if self.training and feature_noise > 0:\n            x = x + torch.randn_like(x) * float(feature_noise)\n\n        x = x * mask.unsqueeze(-1)\n        denom = mask.sum(dim=1, keepdim=True).clamp(min=1).unsqueeze(-1)\n        mean = x.sum(dim=1, keepdim=True) / denom\n\n        query = self.query.unsqueeze(0).expand(len(x), -1, -1)\n        attended, _ = self.attn(\n            query,\n            x,\n            x,\n            key_padding_mask=~mask,\n            need_weights=False,\n        )\n\n        global_target = mean.expand(-1, self.n_target, -1)\n        fused = torch.cat(\n            [\n                attended,\n                query,\n                global_target,\n                torch.abs(attended - global_target),\n                attended * global_target,\n            ],\n            dim=-1,\n        )\n        fused = self.fuse(fused)\n        logits = (\n            fused * self.target_weight.unsqueeze(0)\n        ).sum(dim=-1) + self.target_bias.unsqueeze(0)\n        return logits\n\n\ndef _macro_auc(y_true, pred):\n    aucs = []\n    for j in range(len(LABELS)):\n        y = y_true[:, j]\n        p = pred[:, j]\n        if len(np.unique(y)) == 2:\n            aucs.append(roc_auc_score(y, p))\n    return float(np.mean(aucs)) if aucs else float(\"nan\")\n\n\ndef _ema_update(ema, model, decay=0.995):\n    with torch.no_grad():\n        ema_state = dict(ema.named_parameters())\n        for name, parameter in model.named_parameters():\n            ema_state[name].mul_(decay).add_(\n                parameter.detach(),\n                alpha=1.0 - decay,\n            )\n        ema_buffers = dict(ema.named_buffers())\n        for name, buffer in model.named_buffers():\n            ema_buffers[name].copy_(buffer)\n\n\ndef _batch_rank_loss(logits, labels, gold_rows):\n    losses = []\n    for target in range(logits.shape[1]):\n        y = labels[:, target]\n        positive = torch.nonzero(gold_rows & (y > 0.5), as_tuple=False).flatten()\n        negative = torch.nonzero(gold_rows & (y <= 0.5), as_tuple=False).flatten()\n        if len(positive) == 0 or len(negative) == 0:\n            continue\n\n        n = min(24, len(positive), len(negative))\n        pos = positive[torch.randperm(len(positive), device=logits.device)[:n]]\n        neg = negative[torch.randperm(len(negative), device=logits.device)[:n]]\n        losses.append(\n            F.softplus(\n                -(logits[pos, target] - logits[neg, target])\n            ).mean()\n        )\n    if not losses:\n        return logits.new_tensor(0.0)\n    return torch.stack(losses).mean()\n\n\ndef _predict_head(model, indices, device, batch=512):\n    model.eval()\n    output = []\n    with torch.no_grad():\n        for start in range(0, len(indices), batch):\n            idx = indices[start:start + batch]\n            feat = torch.from_numpy(FEATURES[idx]).to(device)\n            mask = torch.from_numpy(SLOT_MASK[idx]).to(device)\n            logits = model(feat, mask)\n            output.append(torch.sigmoid(logits).cpu().numpy())\n    return np.concatenate(output, axis=0)\n\n\ndef _train_one_fold(fold_id, device):\n    valid = GOLD & (FOLD == int(fold_id))\n    if not valid.any():\n        raise RuntimeError(f\"fold {fold_id} has no gold validation rows\")\n\n    # Held-out gold is excluded even as pseudo-label supervision.\n    allowed = ~valid\n    gold_train = GOLD & allowed\n\n    feature_dim = FEATURES.shape[-1]\n    model = SlotAttentionHead(feature_dim).to(device)\n    ema = copy.deepcopy(model).to(device).eval()\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=7e-4,\n        weight_decay=2e-3,\n    )\n\n    weak_conf = (\n        0.30\n        + 1.70 * np.abs(Y - 0.5) * 2.0\n    ).clip(0.30, 2.0).astype(np.float32)\n\n    target_weight = weak_conf.copy()\n    target_weight[gold_train] = 14.0\n    target_weight[valid] = 0.0\n\n    base_pool = np.flatnonzero(allowed)\n    gold_pool = np.flatnonzero(gold_train)\n\n    # 32 weak-label epochs are cheap because the DINOv3 encoder is frozen.\n    for epoch in range(32):\n        repeated_gold = np.repeat(gold_pool, 8)\n        pool = np.concatenate([base_pool, repeated_gold])\n        rng = np.random.default_rng(SEED + 101 * fold_id + epoch)\n        rng.shuffle(pool)\n\n        model.train()\n        running = 0.0\n        steps = 0\n\n        for start in range(0, len(pool), 256):\n            idx = pool[start:start + 256]\n            if len(idx) < 16:\n                continue\n\n            feat = torch.from_numpy(FEATURES[idx]).to(device)\n            mask = torch.from_numpy(SLOT_MASK[idx]).to(device)\n            labels = torch.from_numpy(Y[idx]).to(device)\n            weights = torch.from_numpy(target_weight[idx]).to(device)\n            gold_rows = torch.from_numpy(GOLD[idx]).to(device)\n\n            optimizer.zero_grad(set_to_none=True)\n            logits = model(\n                feat,\n                mask,\n                modality_dropout=0.13,\n                feature_noise=0.012,\n            )\n\n            bce = F.binary_cross_entropy_with_logits(\n                logits,\n                labels,\n                reduction=\"none\",\n            )\n            loss = (bce * weights).sum() / weights.sum().clamp(min=1.0)\n            loss = loss + 0.055 * _batch_rank_loss(\n                logits,\n                labels,\n                gold_rows,\n            )\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 2.0)\n            optimizer.step()\n            _ema_update(ema, model, 0.995)\n\n            running += float(loss.detach().cpu())\n            steps += 1\n\n        if (epoch + 1) % 8 == 0 or epoch == 0:\n            pred = _predict_head(\n                ema,\n                np.flatnonzero(valid),\n                device,\n            )\n            score = _macro_auc(\n                Y[valid],\n                pred,\n            )\n            print(\n                f\"fold={fold_id} weak epoch={epoch+1:02d}/32 \"\n                f\"loss={running/max(steps,1):.4f} gold_auc={score:.5f}\"\n            )\n\n    # Exact-only refinement: small LR, stronger dropout, no held-out fold.\n    exact_idx = np.flatnonzero(gold_train)\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=1.5e-4,\n        weight_decay=3e-3,\n    )\n\n    for epoch in range(12):\n        rng = np.random.default_rng(SEED + 9000 + 31 * fold_id + epoch)\n        pool = np.tile(exact_idx, 10)\n        rng.shuffle(pool)\n\n        model.train()\n        for start in range(0, len(pool), 128):\n            idx = pool[start:start + 128]\n            feat = torch.from_numpy(FEATURES[idx]).to(device)\n            mask = torch.from_numpy(SLOT_MASK[idx]).to(device)\n            labels = torch.from_numpy(Y[idx]).to(device)\n\n            optimizer.zero_grad(set_to_none=True)\n            logits = model(\n                feat,\n                mask,\n                modality_dropout=0.09,\n                feature_noise=0.008,\n            )\n            loss = F.binary_cross_entropy_with_logits(\n                logits,\n                labels,\n            ) + 0.10 * _batch_rank_loss(\n                logits,\n                labels,\n                torch.ones(len(idx), dtype=torch.bool, device=device),\n            )\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.5)\n            optimizer.step()\n            _ema_update(ema, model, 0.997)\n\n    valid_idx = np.flatnonzero(valid)\n    pred = _predict_head(ema, valid_idx, device)\n\n    checkpoint = {\n        \"version\": 16,\n        \"fold\": int(fold_id),\n        \"labels\": list(LABELS),\n        \"feature_dim\": int(feature_dim),\n        \"hidden\": 384,\n        \"n_slot\": int(N_SLOT),\n        \"state_dict\": {\n            key: value.detach().cpu()\n            for key, value in ema.state_dict().items()\n        },\n    }\n\n    path = OUT / f\"v16_slothead_f{fold_id}.pt\"\n    torch.save(checkpoint, path)\n    return valid_idx, pred, path\n\n\nunique_folds = sorted(\n    int(value)\n    for value in np.unique(FOLD[GOLD])\n    if int(value) >= 0\n)\nif len(unique_folds) < 4:\n    raise RuntimeError(\"V16 expected at least four gold folds\")\n\nOOF_HEAD = np.full_like(Y, np.nan, dtype=np.float32)\nsaved = []\n\n# The head is small; a single GPU is enough after feature extraction.\nDEV_HEAD = torch.device(\n    \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n)\n\nfor fold_id in unique_folds:\n    valid_idx, pred, path = _train_one_fold(\n        fold_id,\n        DEV_HEAD,\n    )\n    OOF_HEAD[valid_idx] = pred\n    saved.append(str(path))\n\ngold_idx = np.flatnonzero(GOLD)\nif not np.isfinite(OOF_HEAD[gold_idx]).all():\n    raise RuntimeError(\"V16 gold OOF is incomplete\")\n\nprint(\"V16 raw gold macro AUC:\", _macro_auc(Y[GOLD], OOF_HEAD[GOLD]))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"142c9d7f-ccad-43a4-b759-5113c2444ee0","cell_type":"code","source":"# -------------------- nested blend selection + dataset manifest --------------------\n\ndef _rank_vector(values):\n    return pd.Series(values).rank(method=\"average\", pct=True).to_numpy(np.float64)\n\ndef _auc(y, p):\n    if len(np.unique(y)) < 2:\n        return np.nan\n    return float(roc_auc_score(y, p))\n\ndef _nested_weight(truth, anchor, specialist, folds, target_index):\n    grid = [0.00, 0.04, 0.08, 0.12, 0.16, 0.20, 0.24, 0.28]\n    nested = np.full(len(truth), np.nan, np.float64)\n\n    for outer in sorted(np.unique(folds)):\n        choose = folds != outer\n        valid = folds == outer\n        best = 0.0\n        best_obj = -1e9\n\n        for weight in grid:\n            pred = (1.0 - weight) * anchor[choose] + weight * specialist[choose]\n            auc = _auc(truth[choose], pred)\n            if not np.isfinite(auc):\n                continue\n            obj = auc - 0.0020 * weight\n            if obj > best_obj:\n                best_obj = obj\n                best = weight\n\n        nested[valid] = (1.0 - best) * anchor[valid] + best * specialist[valid]\n\n    if not np.isfinite(nested).all():\n        return 0.0, 0.0, 0.0\n\n    base_auc = _auc(truth, anchor)\n    nested_auc = _auc(truth, nested)\n    gain = nested_auc - base_auc\n\n    positive = np.flatnonzero(truth > 0.5)\n    negative = np.flatnonzero(truth <= 0.5)\n    rng = np.random.default_rng(16000 + target_index)\n    gains = []\n\n    if len(positive) >= 2 and len(negative) >= 2:\n        for _ in range(320):\n            idx = np.concatenate(\n                [\n                    rng.choice(positive, len(positive), replace=True),\n                    rng.choice(negative, len(negative), replace=True),\n                ]\n            )\n            a0 = _auc(truth[idx], anchor[idx])\n            a1 = _auc(truth[idx], nested[idx])\n            if np.isfinite(a0) and np.isfinite(a1):\n                gains.append(a1 - a0)\n\n    support = float((np.asarray(gains) > 0).mean()) if gains else 0.0\n\n    # Final deployment weight is chosen on all gold only after nested validation.\n    best_weight = 0.0\n    best_obj = base_auc\n    for weight in grid:\n        pred = (1.0 - weight) * anchor + weight * specialist\n        auc = _auc(truth, pred)\n        obj = auc - 0.0020 * weight\n        if obj > best_obj:\n            best_obj = obj\n            best_weight = weight\n\n    protected = {\n        \"ACL\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Baker's\",\n        \"Contusion\",\n    }\n    cap = 0.10 if LABELS[target_index] in protected else 0.22\n\n    if gain >= 0.004 and support >= 0.62:\n        deploy = min(best_weight, cap)\n    elif gain >= 0.002 and support >= 0.59:\n        deploy = min(best_weight, cap * 0.5)\n    else:\n        deploy = 0.0\n\n    return float(deploy), float(gain), float(support)\n\n\nhead_gold = np.zeros((len(gold_idx), len(LABELS)), np.float64)\nfor j in range(len(LABELS)):\n    head_gold[:, j] = _rank_vector(OOF_HEAD[gold_idx, j])\n\nanchor_gold = APPROX_ANCHOR[gold_idx]\ntruth_gold = Y[gold_idx]\nfold_gold = FOLD[gold_idx]\n\ndeploy_weights = {}\nvalidation = {}\n\nfor j, target in enumerate(LABELS):\n    weight, gain, support = _nested_weight(\n        truth_gold[:, j],\n        anchor_gold[:, j],\n        head_gold[:, j],\n        fold_gold,\n        j,\n    )\n    deploy_weights[target] = weight\n    validation[target] = {\n        \"nested_gain\": gain,\n        \"bootstrap_support\": support,\n        \"deploy_weight\": weight,\n        \"head_auc\": _auc(truth_gold[:, j], head_gold[:, j]),\n        \"anchor_auc\": _auc(truth_gold[:, j], anchor_gold[:, j]),\n    }\n\nmanifest = {\n    \"artifact_type\": \"rsna-knee-dinov3-slothead-v16\",\n    \"version\": 16,\n    \"labels\": list(LABELS),\n    \"slots\": [list(slot) for slot in SLOTS],\n    \"image_size\": int(SIZE),\n    \"n_slice\": int(N_SLICE),\n    \"crop_mm\": float(CROP_MM),\n    \"feature\": \"per-slot CLS+mean+central-patch-mean from existing DINOv3 encoder\",\n    \"encoder_checkpoint\": encoder_ckpt.name,\n    \"encoder_cfg\": encoder_cfg,\n    \"encoder_checkpoint_sha256\": hashlib.sha256(\n        encoder_ckpt.read_bytes()\n    ).hexdigest(),\n    \"folds\": [int(v) for v in unique_folds],\n    \"deploy_weights\": deploy_weights,\n    \"inference_shrink\": 0.85,\n    \"validation\": validation,\n    \"oof_source\": OOF_SOURCE,\n}\n\n(OUT / \"v16_manifest.json\").write_text(\n    json.dumps(manifest, indent=2),\n    encoding=\"utf-8\",\n)\n\noof_frame = train_df[[\"StudyInstanceUID\"]].copy()\nfor j, target in enumerate(LABELS):\n    oof_frame[target] = OOF_HEAD[:, j]\noof_frame[\"gold\"] = GOLD.astype(np.uint8)\noof_frame[\"fold\"] = FOLD\noof_frame.to_csv(\n    OUT / \"v16_slothead_oof.csv\",\n    index=False,\n)\n\n# Remove the large train feature cache from the final weights artifact.\n# It is useful during training but is not needed by inference.\ntry:\n    FEATURE_CACHE.unlink()\nexcept OSError:\n    pass\n\nimport zipfile\nzip_path = Path(\"/kaggle/working/rsna-knee-dinov3-slothead-v16.zip\")\nwith zipfile.ZipFile(zip_path, \"w\", compression=zipfile.ZIP_DEFLATED) as archive:\n    for path in sorted(OUT.iterdir()):\n        if path.is_file():\n            archive.write(path, arcname=path.name)\n\nprint(\"saved weights:\", OUT)\nprint(\"zip:\", zip_path)\nprint(\"deploy weights:\")\nfor target in LABELS:\n    info = validation[target]\n    print(\n        f\"{target:18s} weight={deploy_weights[target]:.3f} \"\n        f\"gain={info['nested_gain']:+.4f} \"\n        f\"support={info['bootstrap_support']:.3f}\"\n    )\n","metadata":{},"outputs":[],"execution_count":null},{"id":"196a2a36-1ac1-4b9a-9813-d3ba04a1e8c5","cell_type":"code","source":"# V5.5 TRAIN output contract check\nfrom pathlib import Path as _V55Path\nimport json as _v55_json\n\n_V55_OUT = _V55Path('/kaggle/working/rsna-knee-dinov3-slothead-v16')\n_V55_MANIFEST = _V55_OUT / 'v16_manifest.json'\n_V55_CKPTS = sorted(_V55_OUT.glob('v16_slothead_f*.pt'))\nif not _V55_MANIFEST.is_file():\n    raise RuntimeError('TRAIN failed: v16_manifest.json was not produced')\nif len(_V55_CKPTS) < 4:\n    raise RuntimeError(f'TRAIN failed: only {len(_V55_CKPTS)} slot-head checkpoints were produced')\n_V55_META = _v55_json.loads(_V55_MANIFEST.read_text(encoding='utf-8'))\nif _V55_META.get('artifact_type') != 'rsna-knee-dinov3-slothead-v16':\n    raise RuntimeError('TRAIN failed: artifact contract mismatch')\n(_V55_OUT / 'V55_TRAIN_OUTPUT_READY.txt').write_text(\n    'Attach this notebook Output to RSNA_Knee_DINOsaur_V4_INFER_GoldOOFSlotHead.ipynb\\n',\n    encoding='utf-8',\n)\nprint('[V5.5 TRAIN] OUTPUT READY')\nprint('folder:', _V55_OUT)\nprint('manifest:', _V55_MANIFEST)\nprint('checkpoints:', len(_V55_CKPTS))\nprint('deploy weights:')\nfor _t, _w in _V55_META.get('deploy_weights', {}).items():\n    if float(_w) > 0:\n        print(f'  {_t}: {float(_w):.3f}')\n","metadata":{},"outputs":[],"execution_count":null}]}