{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Multimodal AI for Knee MRI Abnormality Detection\n\n## 1. Overview & Mathematical Formulation\n\nThe objective is to predict the probability of 12 clinical knee abnormalities for each MRI study. The evaluation metric is the macro-averaged Area Under the Receiver Operating Characteristic Curve (AUC ROC).\n\n### 1.1 Evaluation Metric\nLet $N$ be the number of studies and $K=12$ be the number of targets. Let $y_{i,k} \\in \\{0,1\\}$ be the ground truth and $\\hat{p}_{i,k} \\in [0,1]$ be the predicted probability for study $i$ and target $k$. The score is defined as:\n\n$$ \\text{Score} = \\frac{1}{K} \\sum_{k=1}^{K} \\text{AUC}_k(\\mathbf{y}_k, \\mathbf{\\hat{p}}_k) $$\n\nBecause AUC depends only on the rank of the predictions, calibration is irrelevant. However, since the training set lacks complete explicit labels, we derive pseudo-labels from radiology reports, which introduces noise. To handle this, we use a **Confidence-Weighted Binary Cross-Entropy Loss**.\n\n### 1.2 Confidence-Weighted Loss Function\nLet $w_{i,k} \\in [0, 3.0]$ be the confidence weight for target $k$ of study $i$. The loss function is:\n\n$$ \\mathcal{L} = - \\frac{1}{N} \\sum_{i=1}^{N} \\sum_{k=1}^{K} w_{i,k} \\left[ y_{i,k} \\log(\\hat{p}_{i,k}) + (1 - y_{i,k}) \\log(1 - \\hat{p}_{i,k}) \\right] $$\n\nWhere $w_{i,k} = 3.0$ for gold-standard annotated studies, and $w_{i,k} = 0.25 + 0.75 \\cdot c_{i,k}$ for text-derived labels, with $c_{i,k} \\in [0,1]$ being the semantic confidence of the extraction.\n\n### 1.3 Physical Scale Normalization\nMR pixel spacing $s$ (mm/pixel) varies across scanners. Resizing directly to $P \\times P$ distorts physical scale. We crop to a constant physical extent $L = 160$ mm before resampling:\n\n$$ n_{\\text{crop}} = \\left\\lfloor \\frac{L}{s} \\right\\rceil, \\quad s_{\\text{eff}} = \\frac{L}{P} \\approx 0.71 \\text{ mm/pixel} $$\n\n### 1.4 Slot Attention Mechanism\nLet $x_s \\in \\mathbb{R}^d$ be the embedding of slot $s$ (e.g., Sagittal T2). We project it and add a learned positional embedding $e_s$:\n\n$$ h_s = \\phi(x_s) + e_s $$\n\nFor each diagnosis $o$, we compute attention weights over the $S$ available slots:\n\n$$ \\alpha_{o,s} = \\frac{\\exp\\left(\\frac{\\langle h_s, q_o \\rangle}{\\sqrt{d}}\\right) \\cdot m_s}{\\sum_{s'} \\exp\\left(\\frac{\\langle h_{s'}, q_o \\rangle}{\\sqrt{d}}\\right) \\cdot m_{s'}} $$\n\nWhere $q_o$ is the query vector for diagnosis $o$, and $m_s \\in \\{0,1\\}$ is the presence mask. The final context vector $c_o$ and logit $\\ell_o$ are:\n\n$$ c_o = \\sum_{s} \\alpha_{o,s} h_s, \\quad \\ell_o = \\langle c_o, w_o \\rangle + b_o $$\n\n---\n\n## 2. Full Implementation Code","metadata":{}},{"cell_type":"code","source":"# =========================================================\n# RSNA Knee Abnormality Detection: Master Multi-Arm Pipeline (0.936 SOTA)\n# =========================================================\nimport os, sys, glob, time, json, gc, re, hashlib, base64, zlib, contextlib, threading, traceback, warnings\nfrom concurrent.futures import ThreadPoolExecutor, ProcessPoolExecutor, as_completed\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport cv2\nimport timm\nfrom torchvision.models import resnet50\n\nwarnings.filterwarnings('ignore')\ncv2.setNumThreads(1)\nos.environ.setdefault(\"HF_HUB_OFFLINE\", \"1\")\nos.environ.setdefault(\"TRANSFORMERS_OFFLINE\", \"1\")\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\n\nTARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nT0 = time.time()\ndef log(msg): print(f'[{time.time() - T0:7.1f}s] {msg}', flush=True)\n\ndef _comp_root():\n    for _c in (\"/kaggle/input/competitions/rsna-knee-abnormality-detection\", \"/kaggle/input/rsna-knee-abnormality-detection\"):\n        if os.path.isdir(_c): return _c\n    raise RuntimeError(\"competition data not found\")\n_COMP_ROOT = _comp_root()\n\nASSET = Path('/kaggle/input/rsna-knee-bend-dinov3-0917-repro-assets')\nROOT = Path(_COMP_ROOT)\nDINO = Path('/kaggle/input/models/metaresearch/dinov2/pytorch/small/1')\nDEVS = [torch.device(f'cuda:{i}') for i in range(torch.cuda.device_count())]\nSEED = 2026\n\n# =========================================================\n# Stage 1: DINOv2 ViT-Small Foundation Ensemble (0.899)\n# =========================================================\nCROP_MM = 130.0\nCACHE_IMG = 336\nIMG = CACHE_IMG\nGROUP = 3\nN_GROUP_MAX = 1\nCACHE_FRACTION = 0.45\nCACHE_BUDGET_MAX_GB = 24.0\nCACHE_BUDGET_GB = 12.0\nTEST_SHARE = 0.3\nHDR_THREADS = 16\nPIX_THREADS = 12\nORDER_THREADS = 32\nORDER_BUDGET_S = 5400\nAUG_ROT_DEG = 8.0\nAUG_SCALE = 0.08\nAUG_SHIFT = 0.05\nAUG_INTENSITY = 0.1\nLAT_MIN_OFFSET_MM = 20.0\nSLICE_BAND = (0.2, 0.8)\nRULES_NATIVE = {'order': 'normal', 'lat': 'centre', 'slot_fallback': False, 'decode_fill': 'nearest'}\nRULES_LEGACY = {'order': 'dominant_axis', 'lat': 'corner_x', 'slot_fallback': True, 'decode_fill': 'zero'}\nRULES = dict(RULES_NATIVE)\nLEGACY_LAT_OFFSET_MM = 5.0\nEVAL_BATCH = 8\nTIME_BUDGET = 8.0 * 3600\nSLOTS_RECOVERED = [('SAG_FLUID_FS', 'Sagittal', True, True), ('COR_FLUID_FS', 'Coronal', True, True), ('AX_FLUID_FS', 'Axial', True, True), ('SAG_FLUID_NOFS', 'Sagittal', True, False), ('COR_T1', 'Coronal', False, False), ('SAG_T1', 'Sagittal', False, False)]\nSLOTS_PUBLIC = [('SAG_FLUID', 'Sagittal', None, True), ('COR_FLUID', 'Coronal', None, True), ('AX_FLUID', 'Axial', None, True), ('SAG_STRUCT', 'Sagittal', None, False), ('COR_STRUCT', 'Coronal', None, False), ('AX_STRUCT', 'Axial', None, False)]\nSLOT_SCHEME = os.environ.get('SLOT_SCHEME', 'recovered')\nSLOTS = SLOTS_PUBLIC if SLOT_SCHEME == 'public' else SLOTS_RECOVERED\nN_SLOT = len(SLOTS)\nPOOL_PARTS = {'cls_mean': 2, 'cls_mean_focal': 3}\nSLOT_PRIOR_TABLE = {'ACL': (0, 3, 5), 'MCL': (1, 4), 'Medial Meniscus': (0, 1, 3, 4), 'Lateral Meniscus': (0, 1, 3, 4), 'Medial OA': (1, 4, 5), 'Lateral OA': (1, 4, 5), 'PF OA': (0, 2, 5), 'Effusion': (0, 2), 'Synovitis': (0, 2), \"Baker's\": (0,), 'Contusion': (0, 1, 2), 'Fracture': (0, 1, 2, 4, 5)}\nSLOT_PRIOR_STRENGTH = 0.55\nFATSAT_OPTS = {'FS', 'FATSAT', 'FAT_SAT', 'FSAT'}\n_SEP = re.compile('[_\\\\-.]')\n_FATSAT_RX = re.compile('\\\\bfs\\\\b|fatsat|fat sat|\\\\bstir\\\\b|\\\\bspair\\\\b|\\\\bspir\\\\b|\\\\bwe\\\\b|water excit|\\\\btirm\\\\b|\\\\bsting\\\\b|\\\\bfatsup\\\\b')\n_T1_RX = re.compile('\\\\bt1\\\\b|\\\\bt1w\\\\b')\n_T2_RX = re.compile('\\\\bt2\\\\b|\\\\bt2w\\\\b')\n_PD_RX = re.compile('\\\\bpd\\\\b|\\\\bpdw\\\\b|proton|\\\\bdp\\\\b|dens')\n\ndef available_gb():\n    try:\n        with open('/proc/meminfo') as fh:\n            info = {k.strip(): v for k, v in (l.split(':', 1) for l in fh if ':' in l)}\n        return int(info['MemAvailable'].split()[0]) / 1024 ** 2\n    except: return CACHE_BUDGET_GB / CACHE_FRACTION\n\ndef plan_cache(n_study, n_test=0):\n    avail = available_gb()\n    budget = min(avail * CACHE_FRACTION, CACHE_BUDGET_MAX_GB)\n    n_total = n_study + max(n_test, int(TEST_SHARE * n_study))\n    per_slice = n_total * N_SLOT * IMG * IMG\n    afford = int(budget * 1024 ** 3 // max(per_slice, 1))\n    groups = max(1, min(N_GROUP_MAX, afford // GROUP))\n    return groups\n\nN_GROUP = plan_cache(len(pd.read_csv(ROOT / 'train.csv')), len(pd.read_csv(ROOT / 'test.csv')))\nCACHE_SLICES = GROUP * N_GROUP\nHDR_TAGS = ['SeriesDescription', 'SequenceName', 'ScanOptions', 'ScanningSequence', 'RepetitionTime', 'EchoTime', 'Laterality', 'PixelSpacing', 'Rows', 'Columns', 'RescaleSlope', 'RescaleIntercept', 'ImagePositionPatient', 'ImageOrientationPatient']\n\ndef _hdr_vec(s, n):\n    if not isinstance(s, str): return None\n    try: v = [float(x) for x in s.split('|')]\n    except: return None\n    return np.array(v) if len(v) >= n else None\n\ndef side_from_geometry(h):\n    cx = {}\n    for r in h.itertuples(index=False):\n        ipp = _hdr_vec(getattr(r, 'ImagePositionPatient', None), 3)\n        iop = _hdr_vec(getattr(r, 'ImageOrientationPatient', None), 6)\n        ps = _hdr_vec(getattr(r, 'PixelSpacing', None), 2)\n        rows, cols = getattr(r, 'Rows', None), getattr(r, 'Columns', None)\n        if ipp is None or iop is None or ps is None or not rows or not cols: continue\n        try: c = ipp[:3] + iop[:3] * ps[1] * float(cols) / 2 + iop[3:6] * ps[0] * float(rows) / 2\n        except: continue\n        cx.setdefault(r.StudyInstanceUID, []).append(float(c[0]))\n    out = {}\n    for st, xs in cx.items():\n        m = float(np.median(xs))\n        out[st] = None if abs(m) < LAT_MIN_OFFSET_MM else 'R' if m < 0 else 'L'\n    return out\n\ndef side_from_corner_x(h):\n    out = {}\n    for st, g in h.groupby('StudyInstanceUID'):\n        xs = []\n        for r in g.itertuples(index=False):\n            ipp = _hdr_vec(getattr(r, 'ImagePositionPatient', None), 3)\n            if ipp is not None and np.isfinite(ipp).all(): xs.append(float(ipp[0]))\n        if not xs: out[st] = None; continue\n        x = float(np.median(xs))\n        out[st] = None if abs(x) < LEGACY_LAT_OFFSET_MM else 'R' if x < 0 else 'L'\n    return out\n\ndef lat_of(h, tag=''):\n    geo = side_from_corner_x(h) if RULES['lat'] == 'corner_x' else side_from_geometry(h)\n    d = {}\n    for st, g in h.groupby('StudyInstanceUID'):\n        v = [str(x).strip().upper() for x in g['Laterality'].dropna()]\n        if RULES['lat'] == 'corner_x' and 'ImageLaterality' in g.columns:\n            v += [str(x).strip().upper() for x in g['ImageLaterality'].dropna()]\n        v = [x[0] for x in v if x and x[0] in ('L', 'R')]\n        side = v[0] if v else None\n        if side is None: side = geo.get(st)\n        d[st] = side\n    return d\n\ndef probe(item):\n    split, study, series, path = item\n    row = {'split': split, 'StudyInstanceUID': study, 'SeriesInstanceUID': series, 'dir': path}\n    try:\n        files = sorted(e.name for e in os.scandir(path) if e.name.endswith('.dcm'))\n        row['files'] = files; row['n_slices'] = len(files)\n        if not files: return row\n        ds = pydicom.dcmread(os.path.join(path, files[len(files) // 2]), stop_before_pixels=True, force=True)\n        for t in HDR_TAGS:\n            v = getattr(ds, t, None)\n            if v is None: row[t] = None\n            elif isinstance(v, (list, tuple)) or type(v).__name__ == 'MultiValue': row[t] = '|'.join(str(x) for x in v)\n            else: row[t] = str(v)\n    except Exception as exc: row['err'] = str(exc)[:120]\n    return row\n\ndef walk(split):\n    base = ROOT / split\n    items = []\n    if not base.is_dir(): return pd.DataFrame()\n    for study in os.scandir(base):\n        if study.is_dir():\n            for series in os.scandir(study.path):\n                if series.is_dir(): items.append((split, study.name, series.name, series.path))\n    with ThreadPoolExecutor(max_workers=HDR_THREADS) as pool: rows = list(pool.map(probe, items))\n    return pd.DataFrame(rows)\n\ndef annotate(df):\n    desc = df['SeriesDescription'].fillna('') + ' ' + df['SequenceName'].fillna('')\n    desc = desc.str.lower().str.replace(_SEP, ' ', regex=True)\n    opts = df['ScanOptions'].fillna('').str.upper().str.split('|')\n    opts_fs = opts.apply(lambda ts: any(t.strip() in FATSAT_OPTS for t in ts))\n    df['fatsat'] = desc.str.contains(_FATSAT_RX) | opts_fs\n    tr = pd.to_numeric(df['RepetitionTime'], errors='coerce')\n    te = pd.to_numeric(df['EchoTime'], errors='coerce')\n    gre = df['ScanningSequence'].fillna('').str.upper().str.contains('GR')\n    t1, t2, pdw = desc.str.contains(_T1_RX), desc.str.contains(_T2_RX), desc.str.contains(_PD_RX)\n    df['weight'] = np.where(t1 & ~t2 & ~pdw, 'T1', np.where(t2 & ~pdw, 'T2', np.where(pdw, 'PD', np.where(gre, 'GRE', np.where(tr < 800, 'T1', np.where(te > 60, 'T2', np.where(tr >= 800, 'PD', 'UNK')))))))\n    df['fluid'] = np.isin(df['weight'], ['PD', 'T2'])\n    df['px'] = pd.to_numeric(df['PixelSpacing'].fillna('').str.split('|').str[0].replace('', np.nan), errors='coerce')\n    return df\n\ndef pick_slots(series_df, plane_map):\n    series_df = series_df.copy()\n    series_df['plane'] = series_df['SeriesInstanceUID'].map(plane_map)\n    out = {}\n    for study, g in series_df.groupby('StudyInstanceUID'):\n        chosen = {}\n        for name, plane, fluid, fs in SLOTS:\n            sel = (g['plane'] == plane) & (g['fatsat'] == fs)\n            if fluid is not None: sel &= g['fluid'] == fluid\n            cand = g[sel]\n            if len(cand) == 0 and RULES['slot_fallback'] and fluid is False: cand = g[(g['plane'] == plane) & ~g['fatsat']]\n            if len(cand): chosen[name] = cand.sort_values('n_slices', ascending=False).iloc[0]\n        out[study] = chosen\n    return out\n\ndef _natural_key(name): return tuple(int(x) if x.isdigit() else x.lower() for x in re.split('(\\\\d+)', str(name)))\n\ndef _order_dominant_axis(rec):\n    files, d = rec['files'], rec['dir']\n    rows = []\n    for pos, f in enumerate(files):\n        ipp = inst = None\n        try:\n            ds = pydicom.dcmread(os.path.join(d, f), force=True, stop_before_pixels=True, specific_tags=['ImagePositionPatient', 'InstanceNumber'])\n            raw = getattr(ds, 'ImagePositionPatient', None)\n            if raw is not None and len(raw) >= 3:\n                c = np.asarray(raw[:3], dtype=np.float64)\n                if np.isfinite(c).all(): ipp = c\n            n = getattr(ds, 'InstanceNumber', None)\n            if n is not None: inst = float(n)\n        except: pass\n        rows.append((f, ipp, inst, pos))\n    placed = [r for r in rows if r[1] is not None]\n    need = max(2, int(0.8 * len(rows)))\n    if len(placed) >= need:\n        xyz = np.stack([r[1] for r in placed])\n        axis = int(np.argmax(np.ptp(xyz, axis=0)))\n        spare = float(np.nanmedian(xyz[:, axis]))\n        rows.sort(key=lambda r: (float(r[1][axis]) if r[1] is not None else spare, r[2] if r[2] is not None else float('inf'), r[3]))\n    elif sum(r[2] is not None for r in rows) >= need: rows.sort(key=lambda r: (r[2] if r[2] is not None else float('inf'), r[3]))\n    else: rows.sort(key=lambda r: _natural_key(r[0]))\n    return [r[0] for r in rows], True\n\ndef order_slices(rec):\n    if RULES['order'] == 'dominant_axis': return _order_dominant_axis(rec)\n    files, d = rec['files'], rec['dir']\n    keyed = []\n    for f in files:\n        k = None\n        try:\n            ds = pydicom.dcmread(os.path.join(d, f), force=True, stop_before_pixels=True, specific_tags=[(32, 50), (32, 55), (32, 19)])\n            iop = np.asarray(ds.ImageOrientationPatient, dtype=float)\n            ipp = np.asarray(ds.ImagePositionPatient, dtype=float)\n            k = float(np.dot(ipp, np.cross(iop[:3], iop[3:])))\n        except:\n            try: k = float(ds.InstanceNumber)\n            except: pass\n        keyed.append((k, f))\n    if any(k is None for k, _ in keyed): return files, False\n    return [f for _, f in sorted(keyed, key=lambda t: t[0])], True\n\ndef read_slot(rec, n_slice=None, out_size=None):\n    n_slice = GROUP if n_slice is None else n_slice\n    out_size = IMG if out_size is None else out_size\n    files, d, px = rec.get('ordered', []), rec['dir'], rec['px']\n    if not files: return None\n    lo, hi = int(SLICE_BAND[0] * (len(files) - 1)), int(SLICE_BAND[1] * (len(files) - 1))\n    idx = np.unique(np.linspace(lo, hi, n_slice).astype(int)) if hi > lo else np.array([len(files) // 2])\n    while len(idx) < n_slice: idx = np.append(idx, idx[-1])\n    planes = []\n    for i in idx[:n_slice]:\n        try:\n            ds = pydicom.dcmread(os.path.join(d, files[int(i)]), force=True)\n            a = ds.pixel_array.astype(np.float32)\n            a = a * float(getattr(ds, 'RescaleSlope', 1) or 1) + float(getattr(ds, 'RescaleIntercept', 0) or 0)\n        except: a = np.zeros((out_size, out_size), np.float32)\n        planes.append(a)\n    shp = planes[0].shape\n    planes = [p if p.shape == shp else np.zeros(shp, np.float32) for p in planes]\n    vol = np.stack(planes)\n    if px and np.isfinite(px) and px > 0:\n        want = int(round(CROP_MM / px))\n        h, w = shp\n        if 16 < want < min(h, w):\n            cy, cx = h // 2, w // 2\n            half = want // 2\n            vol = vol[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]\n    lo_v, hi_v = np.percentile(vol, [1, 99])\n    vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-06), 0, 1)\n    t = torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0)\n    t = F.interpolate(t, size=(out_size, out_size), mode='bilinear', align_corners=False)\n    return (t.squeeze(0) * 255).round().clamp(0, 255).to(torch.uint8)\n\ndef normalise_laterality(img, plane, lat):\n    if lat != 'R': return img\n    if plane in ('Coronal', 'Axial'): return torch.flip(img, dims=[-1])\n    return torch.flip(img, dims=[0])\n\ndef build_cache(slot_map, plane_map, lat_map, tag):\n    studies = sorted(slot_map)\n    sidx = {s: i for i, s in enumerate(studies)}\n    cache = np.zeros((len(studies), N_SLOT, CACHE_SLICES, IMG, IMG), np.uint8)\n    mask = np.zeros((len(studies), N_SLOT), np.float32)\n    jobs = [(st, k, plane, slot_map[st][name]) for st in studies for k, (name, plane, _, _) in enumerate(SLOTS) if name in slot_map[st]]\n    with ThreadPoolExecutor(max_workers=PIX_THREADS) as pool:\n        for (st, k, plane, _), img in zip(jobs, pool.map(lambda j: read_slot(j[3], CACHE_SLICES, IMG), jobs)):\n            if img is None: continue\n            cache[sidx[st], k] = normalise_laterality(img, plane, lat_map.get(st)).numpy()\n            mask[sidx[st], k] = 1.0\n    gc.collect()\n    return studies, cache, mask\n\nclass SlotHead(nn.Module):\n    def __init__(self, dim, n_slot, n_out, hidden=256, p=0.2, prior=False):\n        super().__init__()\n        self.proj = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU())\n        self.slot_emb = nn.Parameter(torch.randn(n_slot, hidden) * 0.02)\n        self.query = nn.Parameter(torch.randn(n_out, hidden) * 0.02)\n        self.drop = nn.Dropout(p)\n        self.out = nn.Linear(hidden, n_out)\n        self.hidden = hidden\n        p_ = torch.zeros(n_out, n_slot)\n        if prior and n_slot == len(SLOTS) and n_out == len(TARGETS):\n            for t, slots in SLOT_PRIOR_TABLE.items():\n                if t in TARGETS: p_[TARGETS.index(t), list(slots)] = SLOT_PRIOR_STRENGTH\n        self.prior = prior\n        if prior: self.register_buffer('slot_prior', p_)\n    def forward(self, x, mask):\n        h = self.proj(x) + self.slot_emb\n        att = torch.einsum('bsh,oh->bos', h, self.query) / self.hidden ** 0.5\n        if self.prior: att = att + self.slot_prior.unsqueeze(0)\n        att = att.masked_fill(mask.unsqueeze(1) < 0.5, -10000.0).softmax(-1)\n        ctx = self.drop(torch.einsum('bos,bsh->boh', att, h))\n        return (ctx * self.out.weight.unsqueeze(0)).sum(-1) + self.out.bias\n\nclass Model(nn.Module):\n    def __init__(self, backbone, dim, pool='cls_mean', prior=False):\n        super().__init__()\n        self.backbone = backbone\n        self.head = SlotHead(dim * POOL_PARTS[pool], N_SLOT, len(TARGETS), prior=prior)\n        self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n    def forward(self, imgs, mask, img_size=None):\n        B, S = imgs.shape[:2]\n        x = imgs.reshape(B * S, *imgs.shape[2:]).float().div_(255.0)\n        if img_size is not None and img_size != x.shape[-1]: x = F.interpolate(x, size=(img_size, img_size), mode='bilinear', align_corners=False)\n        x = (x - self.mean) / self.std\n        out = self.backbone(pixel_values=x).last_hidden_state\n        feat = torch.cat([out[:, 0], out[:, 1:].mean(1)], dim=1).reshape(B, S, -1)\n        return self.head(feat, mask)\n\ndef find_dinov2(variant='small'):\n    for c in [Path('/kaggle/input/models/metaresearch/dinov2/pytorch/small/1'), Path('/kaggle/input/dinov2-small/pytorch/small/1'), Path('/kaggle/input/dinov2/pytorch/small/1'), Path('/kaggle/input/dinov2-small'), Path('/kaggle/input/dinov2')]:\n        if (c / 'config.json').is_file(): return c\n    for p in Path('/kaggle/input').glob('**/config.json'):\n        if 'dinov2' in str(p).lower(): return p.parent\n    return None\n\n_DINOV2_CACHE = {}\ndef build_model(unfreeze_last, variant='small', pool='cls_mean', prior=False):\n    from transformers import AutoModel\n    p = find_dinov2(variant)\n    if p is None: raise FileNotFoundError('DINOv2 weights not attached')\n    if p not in _DINOV2_CACHE: _DINOV2_CACHE[p] = AutoModel.from_pretrained(str(p))\n    bb = _DINOV2_CACHE[p]\n    for prm in bb.parameters(): prm.requires_grad = False\n    for blk in bb.encoder.layer[max(0, len(bb.encoder.layer) - unfreeze_last):]:\n        for prm in blk.parameters(): prm.requires_grad = True\n    return Model(bb, bb.config.hidden_size, pool=pool, prior=prior)\n\ndef find_weights(name='manifest.json'):\n    for root, dirs, files in os.walk('/kaggle/input'):\n        dirs[:] = [d for d in dirs if d not in ('train_series', 'test_series')]\n        if name not in files: continue\n        try: man = json.loads((Path(root) / name).read_text())\n        except: continue\n        if isinstance(man.get('members'), list) and man['members']: return Path(root)\n    return None\n\ndef _combine(per_member):\n    all_ids = sorted({s for m in per_member for s in m['ids']})\n    pos = {s: i for i, s in enumerate(all_ids)}\n    acc = np.zeros((len(all_ids), len(TARGETS)), np.float64)\n    tot = np.zeros(len(TARGETS), np.float64)\n    for m in per_member:\n        w = np.asarray([float(m.get('weight', 1.0))] * len(TARGETS), dtype=np.float64)\n        r = pd.DataFrame(m['pred']).rank(pct=True).to_numpy()\n        acc[[pos[s] for s in m['ids']]] += r * w[None, :]\n        tot += w\n    return all_ids, acc / tot[None, :]\n\ndef write_submission(pred, studies, test_df, path):\n    sub = pd.DataFrame(pd.DataFrame(pred).rank(pct=True).values, columns=TARGETS)\n    sub.insert(0, 'StudyInstanceUID', studies)\n    sub = test_df[['StudyInstanceUID']].merge(sub, on='StudyInstanceUID', how='left')\n    sub[TARGETS] = sub[TARGETS].fillna(0.5)\n    sub.to_csv(path, index=False)\n    return sub\n\ndef augment(imgs, generator=None):\n    lead = imgs.shape[:-3]\n    x = imgs.reshape(-1, *imgs.shape[-3:]).float()\n    n, dev = x.shape[0], x.device\n    rot = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * (AUG_ROT_DEG * np.pi / 180)\n    sc = 1.0 + torch.rand(n, device=dev, generator=generator) * AUG_SCALE\n    tx = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * AUG_SHIFT\n    ty = (torch.rand(n, device=dev, generator=generator) - 0.5) * 2 * AUG_SHIFT\n    cos, sin = torch.cos(rot) / sc, torch.sin(rot) / sc\n    theta = torch.zeros(n, 2, 3, device=dev, dtype=torch.float32)\n    theta[:, 0, 0], theta[:, 0, 1], theta[:, 0, 2] = cos, -sin, tx\n    theta[:, 1, 0], theta[:, 1, 1], theta[:, 1, 2] = sin, cos, ty\n    grid = F.affine_grid(theta, x.shape, align_corners=False)\n    x = F.grid_sample(x, grid, mode='bilinear', padding_mode='border', align_corners=False)\n    scale = 1.0 + (torch.rand(n, 1, 1, 1, device=dev, generator=generator) - 0.5) * 2 * AUG_INTENSITY\n    x = (x * scale).clamp(0, 255)\n    return x.reshape(*lead, *x.shape[-3:]).to(imgs.dtype)\n\nTTA_OVERLAP = True\nTTA_POOL = 'prob'\nPUBLIC_FRONTIER_TARGET_POOL = {'Fracture': 'max', 'Contusion': 'max', 'Medial Meniscus': 'max', 'Lateral Meniscus': 'max', 'ACL': 'top2', 'MCL': 'top2', \"Baker's\": 'max'}\nTTA_TARGET_POOL = {**PUBLIC_FRONTIER_TARGET_POOL, 'Synovitis': 'original_mean'}\nLEGACY_FOLD_SOFTPOOL_BETA = {'ACL': 6.0, 'MCL': 6.0, 'Medial Meniscus': 8.0, 'Lateral Meniscus': 8.0, \"Baker's\": 8.0, 'Contusion': 8.0, 'Fracture': 10.0}\nLEGACY_FOLD_SOFTPOOL_ALPHA = {'ACL': 0.2, 'MCL': 0.2, 'Medial Meniscus': 0.25, 'Lateral Meniscus': 0.25, \"Baker's\": 0.2, 'Contusion': 0.2, 'Fracture': 0.15}\n\ndef window_starts(n_slice, group, overlap=None):\n    overlap = TTA_OVERLAP if overlap is None else overlap\n    if overlap and n_slice >= group: return list(range(n_slice - group + 1))\n    return [g * group for g in range(max(n_slice // group, 1))]\n\ndef apply_target_window_pool(values, probs, logits, original_probs, mapping, target_idx):\n    for target, mode in mapping.items():\n        j = target_idx[target]\n        if mode == 'max': values[:, j] = probs[:, :, j].max(0).values\n        elif mode == 'mean': values[:, j] = probs[:, :, j].mean(0)\n        elif mode == 'logit_mean': values[:, j] = torch.sigmoid(logits[:, :, j].mean(0))\n        elif mode == 'original_mean': values[:, j] = original_probs[:, :, j].mean(0)\n        elif mode in ('top2', 'top3'):\n            k = min(int(mode[3:]), probs.shape[0])\n            values[:, j] = probs[:, :, j].topk(k, dim=0).values.mean(0)\n    return values\n\ndef legacy_fold_soft_window_pool(original_probs, target_idx):\n    values = original_probs.mean(0).clone()\n    for target, beta in LEGACY_FOLD_SOFTPOOL_BETA.items():\n        j = target_idx[target]\n        x = original_probs[:, :, j]\n        weight = torch.softmax(float(beta) * x, dim=0)\n        values[:, j] = (weight * x).sum(0)\n    return values\n\n@torch.no_grad()\ndef predict_member(model, cache, mask, idx, dev, img_size, group=None, pool=None, starts=None, jitter=False, jitter_seed=SEED, return_public_frontier=False):\n    group = GROUP if group is None else group\n    pool = TTA_POOL if pool is None else pool\n    starts = window_starts(cache.shape[2], group) if starts is None else list(starts)\n    target_idx = {t: j for j, t in enumerate(TARGETS)}\n    jitter_gen = torch.Generator(device=dev)\n    jitter_gen.manual_seed(int(jitter_seed) % (2 ** 63 - 1))\n    model.eval()\n    out, public_frontier_out, public_soft_out = [], [], []\n    for b in range(0, len(idx), EVAL_BATCH):\n        sel = idx[b:b + EVAL_BATCH]\n        m = torch.from_numpy(mask[sel]).to(dev)\n        win_probs, win_logits, win_original_probs = [], [], []\n        for st in starts:\n            rows = torch.from_numpy(np.ascontiguousarray(cache[sel, :, st:st + group])).to(dev)\n            views = [rows] + ([augment(rows, generator=jitter_gen)] if jitter else [])\n            view_probs, view_logits = [], []\n            for view in views:\n                with torch.autocast('cuda', enabled=dev.type == 'cuda'): z = model(view, m, img_size).float()\n                view_logits.append(z)\n                view_probs.append(torch.sigmoid(z))\n            win_logits.append(torch.stack(view_logits).mean(0))\n            win_probs.append(torch.stack(view_probs).mean(0))\n            win_original_probs.append(view_probs[0])\n        probs = torch.stack(win_probs)\n        logits = torch.stack(win_logits)\n        original_probs = torch.stack(win_original_probs)\n        v = torch.sigmoid(logits.mean(0)) if pool == 'logit' else probs.mean(0)\n        v = apply_target_window_pool(v, probs, logits, original_probs, TTA_TARGET_POOL, target_idx)\n        out.append(v.cpu().numpy())\n        if return_public_frontier:\n            public_v = apply_target_window_pool(original_probs.mean(0), original_probs, logits, original_probs, PUBLIC_FRONTIER_TARGET_POOL, target_idx)\n            public_frontier_out.append(public_v.cpu().numpy())\n            public_soft = legacy_fold_soft_window_pool(original_probs, target_idx)\n            public_soft_out.append(public_soft.cpu().numpy())\n    primary = np.concatenate(out) if out else np.zeros((0, len(TARGETS)), np.float32)\n    if not return_public_frontier: return primary\n    public_frontier = np.concatenate(public_frontier_out) if public_frontier_out else np.zeros((0, len(TARGETS)), np.float32)\n    public_soft = np.concatenate(public_soft_out) if public_soft_out else np.zeros((0, len(TARGETS)), np.float32)\n    return primary, public_frontier, public_soft\n\nBUILD_LOCK = threading.Lock()\nSTATE_LOCK = threading.Lock()\n\ndef _run_member(path, m, dev, Cte, Mte, idx, starts, jitter):\n    with BUILD_LOCK:\n        if 'state' in m: state, fp = m['state'], None\n        else:\n            ck = torch.load(Path(path) / m['file'], map_location='cpu', weights_only=False)\n            state, fp = ck['model'], ck.get('fingerprint')\n        model = build_model(int(m['config']['unfreeze_last']), variant=m['config']['variant'], pool=m['config'].get('pool', 'cls_mean'), prior=bool(m['config'].get('prior', False))).to(dev)\n        model.load_state_dict(state)\n    jitter_seed = SEED + int(hashlib.sha256(str(m['id']).encode()).hexdigest()[:8], 16)\n    public_member = 'state' not in m\n    predicted = predict_member(model, Cte, Mte, idx, dev, IMG, starts=starts, jitter=jitter, jitter_seed=jitter_seed, return_public_frontier=public_member)\n    if public_member: p, public_p, public_soft = predicted\n    else: p, public_p, public_soft = predicted, None, None\n    del model, state\n    gc.collect()\n    if dev.type == 'cuda': torch.cuda.empty_cache()\n    return p, public_p, public_soft, (0, 0)\n\ndef combine_public_members_by_fold(per_member, pred_key='pred'):\n    all_ids = sorted({study for member in per_member for study in member['ids']})\n    position = {study: i for i, study in enumerate(all_ids)}\n    groups = {}\n    for i, member in enumerate(per_member):\n        fold = member.get('fold')\n        key = f'fold_{fold}' if fold is not None else f'member_{i}'\n        groups.setdefault(key, []).append(member)\n    fold_ranks, diagnostics = [], []\n    for key, members_in_fold in sorted(groups.items()):\n        matrices = []\n        for member in members_in_fold:\n            values = np.full((len(all_ids), len(TARGETS)), np.nan, np.float64)\n            values[[position[study] for study in member['ids']]] = np.asarray(member[pred_key], np.float64)\n            if np.isnan(values).any():\n                raise RuntimeError(f\"{member.get('id')}: incomplete {pred_key} coverage\")\n            matrices.append(values)\n        raw_fold_mean = np.mean(matrices, axis=0)\n        fold_ranks.append(pd.DataFrame(raw_fold_mean).rank(method='average', pct=True).to_numpy(np.float64))\n        diagnostics.append({'ensemble_group': key, 'members': len(members_in_fold)})\n    if len(fold_ranks) != 5:\n        raise RuntimeError(f'legacy branch requires five folds, found {len(fold_ranks)}')\n    return (all_ids, np.mean(fold_ranks, axis=0), pd.DataFrame(diagnostics))\n\ndef blend_legacy_frontier_and_soft(frontier_rank, soft_rank):\n    output = np.asarray(frontier_rank, np.float64).copy()\n    for j, target in enumerate(TARGETS):\n        alpha = float(LEGACY_FOLD_SOFTPOOL_ALPHA.get(target, 0.0))\n        if alpha: output[:, j] = (1.0 - alpha) * frontier_rank[:, j] + alpha * soft_rank[:, j]\n    return output\n\ndef infer_from_package(path, dev=None):\n    man = json.loads((Path(path) / 'manifest.json').read_text())\n    members = man['members']\n    test_df = pd.read_csv(ROOT / 'test.csv')\n    test_series = pd.read_csv(ROOT / 'test_series.csv')\n    plane_map = dict(zip(test_series['SeriesInstanceUID'], test_series['Anatomical_Plane']))\n    hte = annotate(walk('test_series'))\n    groups = {}\n    for m in members: groups.setdefault(m['pixel_group'], []).append(m)\n    per_member, public_frontier_members = [], []\n    def bank(m, ids, pred, starts, jitter, public_pred=None, public_soft=None):\n        if float(np.std(pred)) < 1e-09: return\n        with STATE_LOCK:\n            per_member.append({'id': m['id'], 'fold': m.get('fold'), 'ids': ids, 'pred': pred, 'weight': m.get('weight', 1.0), 'target_weight': m.get('target_weight'), 'holdout': m.get('holdout')})\n            if public_pred is not None and len(starts) == len(starts_full):\n                public_frontier_members.append({'id': m['id'], 'fold': m.get('fold'), 'ids': ids, 'pred': public_pred, 'soft_pred': public_soft})\n            elif public_pred is not None:\n                log(f\"  {m['id']}: public-frontier vote omitted because only {len(starts)} / {len(starts_full)} windows completed\")\n            all_ids, acc = _combine(per_member)\n            write_submission(acc, all_ids, test_df, 'submission.csv')\n    for gi, (key, gm) in enumerate(groups.items(), 1):\n        cfg = json.loads(key)\n        globals().update(IMG=int(cfg['img']), CACHE_IMG=int(cfg['img']), GROUP=int(cfg['group']), CACHE_SLICES=int(cfg['slices']), N_GROUP=max(int(cfg['slices']) // int(cfg['group']), 1), CROP_MM=float(cfg['crop_mm']), SLICE_BAND=tuple(float(x) for x in cfg['band']), RULES={**RULES_NATIVE, **(cfg.get('rules') or {})} )\n        st_te, Cte, Mte = build_cache(pick_slots(hte, plane_map), plane_map, lat_of(hte, 'test '), f'test g{gi}')\n        idx = np.arange(len(st_te))\n        starts_full = window_starts(Cte.shape[2], GROUP)\n        for m in gm:\n            try:\n                p, public_p, public_soft, _ = _run_member(path, m, dev, Cte, Mte, idx, starts_full, True)\n                bank(m, st_te, p, starts_full, True, public_p, public_soft)\n            except Exception as exc:\n                log(f\"  MEMBER {m['id']} failed: {exc}\")\n    all_ids, acc = _combine(per_member)\n    sub = write_submission(acc, all_ids, test_df, 'submission.csv')\n    if len(public_frontier_members) == len(members):\n        frontier_ids, frontier_acc = _combine(public_frontier_members)\n        write_submission(frontier_acc, frontier_ids, test_df, 'submission_public_0899.csv')\n        fold_ids, fold_frontier, _ = combine_public_members_by_fold(public_frontier_members, 'pred')\n        soft_ids, fold_soft, _ = combine_public_members_by_fold(public_frontier_members, 'soft_pred')\n        legacy_prediction = blend_legacy_frontier_and_soft(fold_frontier, fold_soft)\n        write_submission(legacy_prediction, fold_ids, test_df, 'submission_legacy_fold_blend.csv')\n    return sub\n\ndef run_dinov2():\n    pkg = find_weights()\n    if pkg is None: raise FileNotFoundError('rsna-knee-weights manifest.json not found')\n    infer_from_package(pkg, DEVS[0])\n    public = Path('/kaggle/working/submission_public_0899.csv')\n    if public.is_file(): public.replace('/kaggle/working/submission.csv')\n    for name in ('submission_legacy_fold_blend.csv', 'legacy_fold_diagnostics.csv'):\n        candidate = Path('/kaggle/working') / name\n        if candidate.is_file(): candidate.unlink()\n\nrun_dinov2()\nlog(\"Stage 1 DINOv2 Complete\")\n\n# =========================================================\n# Stage 2: Cross-Series Spatial Attention Arm A5 (0.910)\n# =========================================================\n_A5_SAVED = dict(globals())\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)\n\ndef _find_a5_ckpt():\n    for c in [Path('/kaggle/input/datasets/mattiaangeli/knee-mri-fold-weights'), Path('/kaggle/input/knee-mri-fold-weights'), Path('/kaggle/input/rsna-knee-bend-dinov3-0917-repro-assets/knee-mri-fold-weights')]:\n        if list(c.glob('*_f*.pt')): return c\n    for p in Path('/kaggle/input').glob('**/*_f*.pt'):\n        if 'm_f' in p.name: return p.parent\n    return Path('/kaggle/input')\nCKPT = _find_a5_ckpt()\nDEV = 'cuda' if torch.cuda.is_available() else 'cpu'\nCOMP = Path(_COMP_ROOT)\nSERIES_ROOT = COMP / 'test_series'\nif not SERIES_ROOT.exists(): SERIES_ROOT = COMP / 'train_series'\n\ndef ordered_files(sdir, cap=64):\n    keyed = []\n    for f in Path(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: continue\n        if len(keyed) >= cap * 4: break\n    return [f for _, f in sorted(keyed)]\n\ndef series_side(path):\n    try: return float(pydicom.dcmread(path, stop_before_pixels=True).ImagePositionPatient[0])\n    except: return 0.0\n\ndef read_crop(path):\n    try: ds = pydicom.dcmread(path); arr = ds.pixel_array.astype(np.float32)\n    except: return None\n    try: ps = float(ds.PixelSpacing[0])\n    except: 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    crop = arr[max(0, cy-half):min(arr.shape[0], cy+half), max(0, cx-half):min(arr.shape[1], cx+half)]\n    return crop if crop.size else None\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: return None\n    lo, hi = np.percentile(crop[::4, ::4], [1, 99])\n    return window(crop, lo, hi, flip)\n\ndef build_study_a5(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: continue\n            files = ordered_files(SERIES_ROOT / study / sub.iloc[0].SeriesInstanceUID)\n            if not files: continue\n            flip = plane != 'Sagittal' and series_side(files[0]) < 0\n            lo, hi = SLICE_BAND\n            i0, i1 = int(round(lo * (len(files) - 1))), 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: picks, off = avail, (N_SLICE - len(avail)) // 2\n            for c, p in enumerate(picks):\n                img = render(files[p], flip)\n                if img is None and len(files) > 1: img = render(files[min(len(files) - 1, p + 1)], flip)\n                if img is not None: out[s_i, off + c] = (img * 255).astype(np.uint8)\n            mask[s_i] = len(picks)\n    return idx, out, mask\n\nsub_df = pd.read_csv(COMP / 'sample_submission.csv')\nser_csv = pd.read_csv(COMP / 'test_series.csv')\nif not (COMP / 'test_series').exists(): ser_csv = pd.read_csv(COMP / 'train_series.csv')\nstudies = sub_df.StudyInstanceUID.tolist()\nby = {s: g.to_dict('records') for s, g in ser_csv[ser_csv.StudyInstanceUID.isin(set(studies))].groupby('StudyInstanceUID')}\nN_SLOT_TYPES, MASK_IDX = 6, 0\n\nclass MeanMaxPool(nn.Module):\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) / cnt.clamp(min=1).unsqueeze(1)\n        mx = torch.full((B, D), -10000.0, device=f.device, dtype=f.dtype).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    def __init__(self, d, n_labels=12, n_heads=4, slot_bias=True):\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.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    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: scores = scores + self.slot_bias.t()[slot]\n        idx = sidx.unsqueeze(1).expand(-1, self.k)\n        m = torch.full((B, self.k), float('-inf'), device=scores.device, dtype=scores.dtype).scatter_reduce(0, idx, scores, reduce='amax', include_self=True)\n        e = (scores - m[sidx]).exp()\n        s = torch.zeros(B, self.k, device=scores.device, dtype=scores.dtype).index_add_(0, sidx, e)\n        a = e / s[sidx].clamp(min=1e-06)\n        out = torch.zeros(B, self.k, self.d, device=f.device, dtype=f.dtype).index_add_(0, sidx, a.unsqueeze(-1) * self.val(f).unsqueeze(1))\n        return out, a\n\nclass TokenXAttnPool(nn.Module):\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    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).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    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'): a.num_prefix_tokens = a.num_prefix_tokens + 1\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): x, rope = pos\n        else: x = pos\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): x = blk(x, rope=rope[i])\n            else:\n                for blk in v.blocks: x = blk(x, rope=rope)\n        else: x = v.blocks(x)\n        return v.norm(x)\n    def forward_head(self, x, pre_logits=True): return self.vit.forward_head(x, pre_logits=pre_logits)\n\nclass _GatedDepthBlock(nn.Module):\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    def forward(self, x): return x + self.gamma * self.drop(self.out(self.v(self.norm(x)) * F.silu(self.g(self.norm(x)))))\n\nclass DepthCompress(nn.Module):\n    def __init__(self, n_slice=16, out_ch=3, depth=1, dropout=0.0, ls_init=0.1, imagenet=True):\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([0.485, 0.456, 0.406]).view(1, -1, 1, 1))\n            self.register_buffer('sd', torch.tensor([0.229, 0.224, 0.225]).view(1, -1, 1, 1))\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: z = b(z)\n        z = self.proj(z)\n        if self.imagenet: z = (z - self.mu.to(z.dtype)) / self.sd.to(z.dtype)\n        return z * keep\n\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    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    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    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    def forward(self, x, slot, vmask):\n        T, S, H, W = x.shape\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        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) / cnt.clamp(min=1).unsqueeze(1)\n    mx = torch.full((B, D), -10000.0, device=v.device, dtype=v.dtype).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    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    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    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    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    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    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    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    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    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            self.pool = TokenXAttnPool(d, n_labels) if pool == 'xattn' else LabelAttentionPool(d, n_labels)\n            wd = 3 * d + pe if pool == 'xattn' else 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    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'): 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 NetA5(nn.Module):\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.readout = Readout(pool, D)\n        if cond == 'post': self.slot_emb = nn.Embedding(N_SLOT_TYPES + 1, D, padding_idx=MASK_IDX)\n    def forward(self, im, slot, smeta, sidx, B, vm=None):\n        if self.mixer is not None: im = self.mixer(im, slot, vm)\n        if self.compress is not None: 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: f = f.flatten(1)\n        if self.cond == 'post': f = f + (lambda v: v.unsqueeze(1) if self.tokens else v)(self.slot_emb(slot))\n        return self.readout(f, slot, sidx, B)\n\nmodels_a5 = []\nfor ckpt_path in sorted(CKPT.glob('*_f*.pt')):\n    if 'm_f' not in ckpt_path.name and 'fold' not in str(ckpt_path).lower(): continue\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': enc = ViTSlotToken(enc, N_SLOT_TYPES)\n    m = NetA5(enc, cfg['cond'], cfg.get('n_meta', 0), cfg['pool'], stem=_stem, n_slice=cfg.get('n_slice', 16))\n    m.load_state_dict(z['state_dict'], strict=False)\n    models_a5.append(m.eval())\nmodels_a5 = [m.to(DEV).eval() for m in models_a5]\n\ndef _norm_a5(im, k='none'):\n    if k == '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-06) * m\n    if k == 'imagenet':\n        m = (im > 0).float()\n        return (im - 0.485) / 0.229 * m\n    return im\n\n@torch.no_grad()\ndef _micro_a5(images, masks):\n    ims, slots, sidx, vms = [], [], [], []\n    for b in range(len(masks)):\n        present = np.nonzero(masks[b] > 0)[0]\n        if len(present) == 0: continue\n        blk = images[b][present]\n        ims.append(torch.from_numpy(blk))\n        vms.append(torch.from_numpy(blk.reshape(blk.shape[0], blk.shape[1], -1).max(2) > 0))\n        slots.append(torch.from_numpy(present + 1).long())\n        sidx.append(torch.full((len(present),), b, dtype=torch.long))\n    out = np.full((len(models_a5), len(masks), len(TARGETS)), np.nan, np.float32)\n    if not ims: return out\n    im = _norm_a5(torch.cat(ims).to(DEV).float().div_(255.0), 'none')\n    sl = torch.cat(slots).to(DEV)\n    si = torch.cat(sidx).to(DEV)\n    vm = torch.cat(vms).to(DEV)\n    with torch.autocast('cuda' if str(DEV).startswith('cuda') else 'cpu', dtype=torch.bfloat16, enabled=str(DEV).startswith('cuda')):\n        for fold_index, model in enumerate(models_a5):\n            out[fold_index] = torch.sigmoid(model(im, sl, torch.zeros(len(sl), 0, device=DEV), si, len(masks), vm=vm)).float().cpu().numpy()\n    return out\n\ndef predict_a5(images, masks):\n    out = np.full((len(models_a5), len(masks), len(TARGETS)), np.nan, np.float32)\n    for a in range(0, len(masks), 8):\n        out[:, a:min(a + 8, len(masks))] = _micro_a5(images[a:min(a + 8, len(masks))], masks[a:min(a + 8, len(masks))])\n    return out\n\npreds_a5 = np.full((len(models_a5), len(studies), len(TARGETS)), np.nan, np.float32)\nwith ProcessPoolExecutor(max_workers=4) as ex:\n    futs = [ex.submit(build_study_a5, (i, s, by.get(s, []))) for i, s in enumerate(studies)]\n    imgs = np.zeros((len(studies), N_SLOT, N_SLICE, SIZE, SIZE), np.uint8)\n    msks = np.zeros((len(studies), N_SLOT), np.uint8)\n    for f in as_completed(futs):\n        i, a, k = f.result()\n        imgs[i], msks[i] = a, k\npreds_a5 = predict_a5(imgs, msks)\ndel imgs, msks; gc.collect()\n\nA5_W = 0.45\n_a5_ok = np.isfinite(preds_a5).all(axis=(0, 2))\n_a5_rank_mean = np.zeros((len(studies), len(TARGETS)), np.float64)\nfor fold_index in range(preds_a5.shape[0]):\n    fold = preds_a5[fold_index][_a5_ok]\n    ordinal = fold.argsort(0).argsort(0).astype(np.float64)\n    _a5_rank_mean[_a5_ok] += ordinal / max(len(fold) - 1, 1)\n_a5_rank_mean /= preds_a5.shape[0]\nA5_PREDS = dict(zip(sub_df['StudyInstanceUID'].astype(str), _a5_rank_mean.astype(np.float32)))\nfor _a5k, _a5v in _A5_SAVED.items(): globals()[_a5k] = _a5v\ndel _A5_SAVED, _a5k, _a5v\n_a5_sub = pd.read_csv('/kaggle/working/submission.csv', dtype={'StudyInstanceUID': str})\n_a5_ours = np.stack([A5_PREDS[_u] for _u in _a5_sub['StudyInstanceUID'].astype(str)])\n_a5_base_rank = _a5_sub[TARGETS].rank(method='average', pct=True)\n_a5_ours_rank = pd.DataFrame(_a5_ours, columns=TARGETS, index=_a5_sub.index).rank(method='average', pct=True)\nfor _lbl in TARGETS:\n    _w = {'ACL': 0.54, 'MCL': 0.54, 'Contusion': 0.52, 'Medial Meniscus': 0.48, 'Lateral Meniscus': 0.48, 'PF OA': 0.48, 'Lateral OA': 0.45, 'Medial OA': 0.45, 'Effusion': 0.45, 'Synovitis': 0.45, \"Baker's\": 0.38, 'Fracture': 0.38}.get(_lbl, 0.48)\n    _a5_sub[_lbl] = (1.0 - _w) * _a5_base_rank[_lbl] + _w * _a5_ours_rank[_lbl]\n_a5_sub[TARGETS] = _a5_sub[TARGETS].rank(method='average', pct=True)\n_a5_sub.to_csv('/kaggle/working/submission.csv', index=False)\nlog(\"Stage 2 A5 Complete\")\n\n# =========================================================\n# Stage 3: Dual-Layout RadImageNet ResNet-50 & V18 Calibration (0.920)\n# =========================================================\n_RAD_LABELS = TARGETS\n_RAD_ALPHA = 0.5\n_RAD_EXCLUDE = (\"Baker's\", 'Fracture')\n_RAD_TOKEN_DIM, _RAD_HEAD_DIM = 2048, 512\n\ndef _rad_find_file(name):\n    for p in Path('/kaggle/input').glob(f'**/{name}'): return p\n    return None\n\nclass _RadEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = nn.Sequential(*list(resnet50(weights=None).children())[:-2])\n    def forward(self, image): return self.backbone(image).mean(dim=(2, 3))\n\nclass _RadHead(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.project = nn.Sequential(nn.LayerNorm(_RAD_TOKEN_DIM), nn.Linear(_RAD_TOKEN_DIM, _RAD_HEAD_DIM), nn.GELU())\n        self.plane = nn.Parameter(torch.randn(N_SLOT, _RAD_HEAD_DIM) * 0.01)\n        self.position = nn.Parameter(torch.randn(CACHE_SLICES, _RAD_HEAD_DIM) * 0.01)\n        self.query = nn.Parameter(torch.randn(len(_RAD_LABELS), _RAD_HEAD_DIM) * 0.02)\n        self.attn = nn.MultiheadAttention(_RAD_HEAD_DIM, 8, dropout=0.1, batch_first=True)\n        self.fuse = nn.Sequential(nn.LayerNorm(_RAD_HEAD_DIM * 4), nn.Linear(_RAD_HEAD_DIM * 4, _RAD_HEAD_DIM), nn.GELU(), nn.Dropout(0.15))\n        self.weight = nn.Parameter(torch.randn(len(_RAD_LABELS), _RAD_HEAD_DIM) * 0.02)\n        self.bias = nn.Parameter(torch.zeros(len(_RAD_LABELS)))\n    def forward(self, feature, mask):\n        token = self.project(feature.float())\n        token = token.view(len(token), N_SLOT, CACHE_SLICES, _RAD_HEAD_DIM)\n        token = token + self.plane[None, :, None] + self.position[None, None]\n        token = token.flatten(1, 2)\n        key_padding = mask <= 0\n        if key_padding.all(1).any(): key_padding[key_padding.all(1), 0] = False\n        query = self.query.unsqueeze(0).expand(len(token), -1, -1)\n        attended = query + self.attn(query, token, token, key_padding_mask=key_padding, need_weights=False)[0]\n        mean = (token * mask.unsqueeze(-1)).sum(1, keepdim=True) / mask.sum(1, keepdim=True).clamp_min(1).unsqueeze(-1)\n        fused = self.fuse(torch.cat([attended, mean.expand(-1, len(_RAD_LABELS), -1), torch.abs(attended - mean.expand(-1, len(_RAD_LABELS), -1)), attended * mean.expand(-1, len(_RAD_LABELS), -1)], dim=-1))\n        return (fused * self.weight.unsqueeze(0)).sum(-1) + self.bias\n\ndef _rad_rank_columns(values): return pd.DataFrame(np.asarray(values, dtype=np.float64)).rank(method='average', pct=True).to_numpy(np.float64)\n\ndef _rad_encode(encoder, pixels, slot_mask, device):\n    n, slots, slices, height, width = pixels.shape\n    features = np.zeros((n, slots * slices, _RAD_TOKEN_DIM), np.float16)\n    token_mask = np.repeat(slot_mask[:, :, None], slices, axis=2).reshape(n, -1)\n    valid = np.flatnonzero(token_mask.reshape(-1) > 0)\n    flat = pixels.reshape(-1, height, width)\n    for start in range(0, len(valid), 96):\n        indices = valid[start:start + 96]\n        image = torch.from_numpy(flat[indices]).to(device).float().div_(127.5).sub_(1.0).unsqueeze(1).expand(-1, 3, -1, -1).contiguous()\n        # إصلاح الخطأ: استخدام torch.no_grad() لمنع تتبع التدرجات\n        with torch.no_grad():\n            with torch.autocast('cuda'): feature = encoder(image)\n        features.reshape(-1, _RAD_TOKEN_DIM)[indices] = feature.float().cpu().numpy()\n    return features, token_mask.astype(np.float32)\n\ndef _rad_main():\n    primary = Path('/kaggle/working/submission.csv')\n    test = pd.read_csv(ROOT / 'test.csv', dtype={'StudyInstanceUID': str})\n    expected_ids = test.StudyInstanceUID.astype(str).tolist()\n    baseline = pd.read_csv(primary, dtype={'StudyInstanceUID': str})\n    device = torch.device('cuda:0')\n    test_series = pd.read_csv(ROOT / 'test_series.csv', dtype={'StudyInstanceUID': str, 'SeriesInstanceUID': str})\n    plane = dict(zip(test_series.SeriesInstanceUID, test_series.Anatomical_Plane))\n    \n    enc_path = _rad_find_file('ResNet50.pt')\n    if not enc_path: log(\"RadImageNet Encoder not found. Skipping Stage 3.\"); return\n    encoder = _RadEncoder()\n    encoder.load_state_dict(torch.load(enc_path, map_location='cpu', weights_only=True), strict=True)\n    encoder.eval().to(device)\n    for p in encoder.parameters(): p.requires_grad_(False)\n    \n    def cache(slots, crop, tag):\n        globals().update(SLOTS=list(slots), N_SLOT=len(slots), CACHE_SLICES=8, IMG=224, CACHE_IMG=224, CROP_MM=float(crop), RULES=dict(RULES_LEGACY))\n        headers = annotate(walk('test_series'))\n        studies, pixels, masks = build_cache(pick_slots(headers, plane), plane, lat_of(headers, tag + ' '), tag)\n        positions = {str(uid): index for index, uid in enumerate(studies)}\n        order = np.asarray([positions[uid] for uid in expected_ids], dtype=np.int64)\n        return pixels[order], masks[order]\n        \n    public_slots = [('SAG_FS', 'Sagittal', None, True), ('COR_FS', 'Coronal', None, True), ('AX_FS', 'Axial', None, True)]\n    pixels, masks = cache(public_slots, 10000.0, 'test-e10')\n    features, token_mask = _rad_encode(encoder, pixels, masks, device)\n    \n    heads_path = _rad_find_file('v52_radimagenet_heads.pt') or _rad_find_file('rad_head_f0.pt')\n    if not heads_path: log(\"RadImageNet Heads not found. Skipping Stage 3 blend.\"); return\n    heads = []\n    try:\n        payload = torch.load(heads_path, map_location='cpu', weights_only=True)\n        folds = payload.get('folds', [payload])\n        for record in folds:\n            head = _RadHead().to(device).eval()\n            head.load_state_dict(record['state_dict'], strict=True)\n            heads.append(head)\n    except Exception as e:\n        log(f\"Error loading Rad heads: {e}\")\n        return\n    \n    if not heads: log(\"No Rad heads loaded. Skipping Stage 3 blend.\"); return\n    \n    reference_predictions = [torch.sigmoid(head(torch.from_numpy(features).to(device), torch.from_numpy(token_mask.astype(np.float32)).to(device)).float()).cpu().numpy() for head in heads]\n    reference_probability = np.mean(np.stack(reference_predictions), axis=0)\n    reference_rank = _rad_rank_columns(reference_probability)\n    \n    baseline_rank = _rad_rank_columns(baseline[_RAD_LABELS].to_numpy())\n    e10 = baseline.copy()\n    for index, target in enumerate(_RAD_LABELS):\n        if target not in _RAD_EXCLUDE: e10[target] = (1.0 - _RAD_ALPHA) * baseline_rank[:, index] + _RAD_ALPHA * reference_rank[:, index]\n    e10.to_csv(primary, index=False)\n    log(\"Stage 3 RadImageNet Complete\")\n\ntry: _rad_main()\nexcept Exception as e: log(f\"Stage 3 failed: {e}\")\n\n# =========================================================\n# Stage 4 & 5: CoAtNet Raptor Arm & Final Fusion (0.936)\n# =========================================================\nIMG_C = 336\nCROP_MM_C = 140.0\nSLOTS_C = [(\"Sagittal\", 1, 18), (\"Sagittal\", 0, 14), (\"Coronal\", 1, 12), (\"Coronal\", 0, 8), (\"Axial\", -1, 12)]\nMAXS_C = sum(s[2] for s in SLOTS_C)\nK_EVAL_C = 62\n_MEAN_C = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n_STD_C = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\ndef build_backbone(arch, pretrained=False):\n    hybrid = arch.startswith((\"maxvit\", \"maxxvit\", \"coatnet\", \"coat_\", \"convnext\"))\n    is_vit = (not hybrid) and any(k in arch for k in (\"vit\", \"deit\", \"dinov2\", \"eva\", \"beit\"))\n    kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n    if is_vit: kw.update(global_pool=\"token\", dynamic_img_size=True)\n    else: kw.update(global_pool=\"avg\")\n    return timm.create_model(arch, **kw)\n\nclass RaptorClassifier(nn.Module):\n    def __init__(self, backbone, F_dim=768, n=12, drop=0.2):\n        super().__init__()\n        self.backbone = backbone\n        self.norm = nn.LayerNorm(F_dim)\n        self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop), nn.Linear(256, n))\n        self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n        self.clsb = nn.Parameter(torch.zeros(n))\n        nn.init.trunc_normal_(self.clsW, std=0.02)\n    def encode(self, x): return self.backbone(x.flatten(0, 1)).view(x.shape[0], x.shape[1], -1)\n    def head(self, feats):\n        h = self.norm(feats)\n        a = torch.softmax(self.att(h), dim=1)\n        return (torch.einsum(\"bkn,bkf->bnf\", a, h) * self.clsW).sum(-1) + self.clsb\n    def forward(self, x): return self.head(self.encode(x))\n\ndef find_weight_file(fname):\n    for d in [\"/kaggle/input/raptor-knee-arms\", \"/kaggle/input/raptor-knee-widedense\", \"/kaggle/input/raptor-cnn336\"]:\n        p = Path(d) / fname\n        if p.exists(): return p\n    for p in Path('/kaggle/input').glob(f'**/{fname}'): return p\n    return None\n\ndef _eval_centers_c(mask, D, k):\n    valid = np.where(mask > 0)[0]\n    if len(valid) < 3: valid = np.arange(min(3, D))\n    lo, hi = int(valid.min()), int(valid.max())\n    cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi] or [max(1, min((lo + hi) // 2, D - 2))]\n    idx = np.linspace(0, len(cs) - 1, k).round().astype(int)\n    return [cs[i] for i in idx]\n\ndef eval_windows_c(vol, mask, k, res):\n    D = vol.shape[0]\n    cs = _eval_centers_c(mask, D, k)\n    wins = np.empty((len(cs), 3, res, res), np.float32)\n    for j, c in enumerate(cs):\n        c = max(1, min(c, D - 2))\n        tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0\n        t = torch.from_numpy(tri)\n        if t.shape[-1] != res: t = F.interpolate(t[None], size=(res, res), mode=\"bilinear\", align_corners=False)[0]\n        wins[j] = ((t - _MEAN_C) / _STD_C).numpy()\n    return torch.from_numpy(wins)\n\ndef read_px_c(f):\n    d = pydicom.dcmread(f)\n    a = d.pixel_array.astype(np.float32)\n    if str(getattr(d, 'PhotometricInterpretation', '')) == 'MONOCHROME1': a = a.max() - a\n    return a\n\ndef mm_crop_resize_c(a, ps):\n    h, w = a.shape; cpx = min(int(round(CROP_MM_C / max(ps, 1e-3))), min(h, w))\n    a = a[(h - cpx) // 2:(h - cpx) // 2 + cpx, (w - cpx) // 2:(w - cpx) // 2 + cpx]\n    return cv2.resize(a, (IMG_C, IMG_C), interpolation=cv2.INTER_AREA)\n\ndef _pick_series_for_slot_c(rows, plane, fluid, used):\n    cands = [r for r in rows if r['Anatomical_Plane'] == plane and r['SeriesInstanceUID'] not in used]\n    if fluid in (0, 1):\n        pref = [r for r in cands if int(r.get('Fluid_Sensitive', 0) or 0) == fluid]\n        if pref: return pref[0]\n    return cands[0] if cands else None\n\ndef build_study_coatnet(sid, ser_records, tsdir):\n    primary_volume = np.zeros((MAXS_C, IMG_C, IMG_C), np.uint8)\n    used, offset = set(), 0\n    for plane, fluid, count in SLOTS_C:\n        record = _pick_series_for_slot_c(ser_records.get(sid, []), plane, fluid, used)\n        if record is None: offset += count; continue\n        used.add(record[\"SeriesInstanceUID\"])\n        sdir = f\"{tsdir}/{sid}/{record['SeriesInstanceUID']}\"\n        files = sorted(glob.glob(sdir + \"/*.dcm\"))\n        if not files: offset += count; continue\n        ps_list = []\n        ordered = []\n        for f in files:\n            try:\n                h = pydicom.dcmread(f, stop_before_pixels=True)\n                iop = getattr(h, 'ImageOrientationPatient', None); ipp = getattr(h, 'ImagePositionPatient', None)\n                if iop and ipp and len(iop) == 6:\n                    n = np.cross(np.array(iop[:3], float), np.array(iop[3:], float))\n                    pos = float(np.dot(np.array(ipp, float), n))\n                else: pos = float(getattr(h, 'InstanceNumber', 0) or 0)\n                ps = getattr(h, 'PixelSpacing', None); ps = float(ps[0]) if ps else 0.5\n                ps_list.append(ps); ordered.append((pos, f, ps))\n            except: ordered.append((0.0, f, 0.5))\n        ordered.sort(key=lambda x: x[0])\n        med_ps = float(np.median(ps_list)) if ps_list else 0.5\n        if len(ordered) > 1:\n            picks = np.linspace(int(0.02 * (len(ordered) - 1)), int(0.98 * (len(ordered) - 1)), count).round().astype(int)\n        else: picks = np.zeros(count, dtype=int)\n        for local_idx, pos in enumerate(picks):\n            if offset + local_idx >= MAXS_C: break\n            pos = min(int(pos), len(ordered) - 1)\n            try: arr = read_px_c(ordered[pos][1])\n            except: continue\n            arr = np.clip((arr - np.percentile(arr, 2)) / (np.percentile(arr, 98) - np.percentile(arr, 2) + 1e-6), 0, 1)\n            primary_volume[offset + local_idx] = (mm_crop_resize_c(arr, ordered[pos][2] or med_ps) * 255).astype(np.uint8)\n        offset += count\n    mask = (primary_volume.reshape(MAXS_C, -1).sum(1) > 0).astype(np.uint8)\n    return primary_volume, mask\n\ndef run_coatnet():\n    tsdir = str(ROOT / \"test_series\")\n    test = pd.read_csv(ROOT / \"test.csv\"); test[\"StudyInstanceUID\"] = test[\"StudyInstanceUID\"].astype(str)\n    test_ids = test[\"StudyInstanceUID\"].tolist()\n    tser = pd.read_csv(ROOT / \"test_series.csv\")\n    tser[\"StudyInstanceUID\"] = tser[\"StudyInstanceUID\"].astype(str)\n    SER = {k: v.to_dict(\"records\") for k, v in tser.groupby(\"StudyInstanceUID\")}\n    \n    swa_path = find_weight_file(\"raptor_ft_coatnet_v5_full_swa.pt\")\n    v4_path = find_weight_file(\"raptor_ft_coatnet_v4_full.pt\")\n    \n    if not swa_path: log(\"CoAtNet SWA weights not found. Skipping Stage 4.\"); return None\n    \n    dev0 = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    ck = torch.load(swa_path, map_location=\"cpu\", weights_only=False)\n    bb = build_backbone(ck.get(\"arch\", \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\"))\n    model = RaptorClassifier(bb, F_dim=bb.num_features).to(dev0).eval()\n    model.load_state_dict(ck[\"model\"], strict=True)\n    res = int(ck.get(\"res\", 384))\n    \n    primary_preds = np.full((len(test_ids), len(TARGETS)), 0.5, np.float32)\n    \n    for i, sid in enumerate(test_ids):\n        vol, mask = build_study_coatnet(sid, SER, tsdir)\n        wins = eval_windows_c(vol, mask, K_EVAL_C, res)\n        with torch.no_grad(), torch.autocast(\"cuda\", dtype=torch.float16):\n            primary_preds[i] = torch.sigmoid(model(wins.unsqueeze(0).to(dev0)).float()[0]).cpu().numpy()\n    \n    if v4_path and torch.cuda.device_count() >= 2:\n        dev1 = torch.device(\"cuda:1\")\n        ck2 = torch.load(v4_path, map_location=\"cpu\", weights_only=False)\n        bb2 = build_backbone(ck2.get(\"arch\", \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\"))\n        model2 = RaptorClassifier(bb2, F_dim=bb2.num_features).to(dev1).eval()\n        model2.load_state_dict(ck2[\"model\"], strict=True)\n        res2 = int(ck2.get(\"res\", 384))\n        legacy_preds = np.full_like(primary_preds, 0.5)\n        for i, sid in enumerate(test_ids):\n            vol, mask = build_study_coatnet(sid, SER, tsdir)\n            wins = eval_windows_c(vol, mask, 42, res2)\n            with torch.no_grad(), torch.autocast(\"cuda\", dtype=torch.float16):\n                legacy_preds[i] = torch.sigmoid(model2(wins.unsqueeze(0).to(dev1)).float()[0]).cpu().numpy()\n        \n        p_rank = pd.DataFrame(primary_preds.clip(0, 1)).rank(pct=True).to_numpy()\n        l_rank = pd.DataFrame(legacy_preds.clip(0, 1)).rank(pct=True).to_numpy()\n        comp_w = {\"MCL\": 0.16, \"Medial OA\": 0.10, \"PF OA\": 0.12, \"Effusion\": 0.10, \"Synovitis\": 0.16, \"Baker's\": 0.12, \"Contusion\": 0.12}\n        for j, t in enumerate(TARGETS):\n            w = comp_w.get(t, 0.0)\n            if w > 0:\n                corr = float(np.corrcoef(p_rank[:, j], l_rank[:, j])[0, 1])\n                if corr > 0.992: w *= 0.5\n                elif corr < 0.65: w *= 0.4\n                p_rank[:, j] = (1.0 - w) * p_rank[:, j] + w * l_rank[:, j]\n        primary_preds = p_rank\n    else:\n        primary_preds = pd.DataFrame(primary_preds.clip(0, 1)).rank(pct=True).to_numpy()\n        \n    sub = pd.DataFrame(primary_preds.astype(np.float32), columns=TARGETS)\n    sub.insert(0, \"StudyInstanceUID\", test_ids)\n    sub.to_csv('/kaggle/working/submission_coatnet.csv', index=False)\n    log(\"Stage 4 CoAtNet Complete\")\n    return sub\n\ntry:\n    coatnet_sub = run_coatnet()\n    if coatnet_sub is not None:\n        tr_sub = pd.read_csv('/kaggle/working/submission.csv', dtype={'StudyInstanceUID': str})\n        cr_sub = pd.read_csv('/kaggle/working/submission_coatnet.csv', dtype={'StudyInstanceUID': str})\n        if tr_sub.columns.tolist() == cr_sub.columns.tolist() and tr_sub['StudyInstanceUID'].tolist() == cr_sub['StudyInstanceUID'].tolist():\n            tr_r = tr_sub[TARGETS].rank(pct=True)\n            cr_r = cr_sub[TARGETS].rank(pct=True)\n            cw = {'Lateral Meniscus': 0.62, 'Fracture': 0.62, 'Medial Meniscus': 0.58, 'Lateral OA': 0.56, 'PF OA': 0.55, 'ACL': 0.54, 'MCL': 0.52, 'Medial OA': 0.52, 'Effusion': 0.50, 'Contusion': 0.48, 'Synovitis': 0.45, \"Baker's\": 0.42}\n            for t in TARGETS:\n                w = cw.get(t, 0.50)\n                tr_sub[t] = (1.0 - w) * tr_r[t] + w * cr_r[t]\n            tr_sub[TARGETS] = tr_sub[TARGETS].rank(pct=True)\n            tr_sub.to_csv('/kaggle/working/submission.csv', index=False)\n            log(\"Final DINOsaur V4.2 Fusion Complete!\")\nexcept Exception as e:\n    log(f\"Stage 4 failed: {e}\")\n\nlog(\"Pipeline Finished Successfully. submission.csv is ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-06T06:34:42.212781Z","iopub.execute_input":"2026-09-06T06:34:42.213164Z","iopub.status.idle":"2026-09-06T06:40:32.582026Z","shell.execute_reply.started":"2026-09-06T06:34:42.213122Z","shell.execute_reply":"2026-09-06T06:40:32.581334Z"}},"outputs":[],"execution_count":null}]}