{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee v5 — RadImageNet ResNet-50 encoder\n\nv2/v3/v4 all started from **ImageNet** weights: a feature vocabulary learned on natural photographs, where\nnothing looks like a torn meniscus on a proton-density sequence. v5 swaps in **RadImageNet ResNet-50**\n(1.35M medical images; MIT-licensed official checkpoint) and keeps what worked in v4 — three plane slots,\ngated slice attention, blended soft labels from two LLM readers, EMA weights.\n\n**One change of method:** checkpoints are selected on the **250-study derived holdout**, not the 58 annotated\nstudies. v4 proved the 58-study set misleads — it ranked v4 *last* (0.847) while the holdout ranked it\nfirst (0.861), and the leaderboard agreed with the holdout (0.883, our best). The 58 are still reported,\njust not trusted for selection."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import os, glob, time, copy\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom sklearn.metrics import roc_auc_score\nfrom concurrent.futures import ThreadPoolExecutor\n\nT0 = time.time()\nPATH = os.path.dirname(glob.glob('/kaggle/input/**/train.csv', recursive=True)[0])\nRADW = glob.glob('/kaggle/input/**/ResNet50.pt', recursive=True)[0]\nDEVICE = 'cuda'\nN_SLICES, CACHE_SIZE, SIZE = 16, 240, 224\nEPOCHS, BATCH, ACCUM = 8, 3, 3\nSLOTS = [('sagfs',  lambda g: (g['Anatomical_Plane']=='Sagittal') & (g['Fluid_Sensitive']==1)),\n         ('cor',    lambda g: (g['Anatomical_Plane']=='Coronal')),\n         ('ax',     lambda g: (g['Anatomical_Plane']=='Axial'))]\ndef log(m): print(f'[{time.time()-T0:7.0f}s] {m}', flush=True)\nlog(f'{torch.cuda.get_device_name(0)} | RadImageNet: {RADW}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"tr = pd.read_csv(os.path.join(PATH, 'train.csv'))\nser = pd.read_csv(os.path.join(PATH, 'train_series.csv'))\nLABELS = list(tr.columns[2:14])\n\nl1 = pd.read_csv(glob.glob('/kaggle/input/**/llm_labels_v4_blend.csv', recursive=True)[0]).set_index('StudyInstanceUID')[LABELS]\nl2 = pd.read_csv(glob.glob('/kaggle/input/**/report_labels_v2.csv', recursive=True)[0]).set_index('StudyInstanceUID')[LABELS]\nl2 = l2.apply(pd.to_numeric, errors='coerce')\ncommon = l1.index.intersection(l2.index)\nsoft = pd.concat([(l1.loc[common] + l2.loc[common]) / 2.0, l1.loc[l1.index.difference(common)]])\n\nannotated = tr[tr[LABELS].notna().all(axis=1)]\nval_ids = set(annotated['StudyInstanceUID'])\ny_true = tr.set_index('StudyInstanceUID').loc[sorted(val_ids), LABELS].astype(float)\n\nleak = tr['Report'].isin(set(annotated['Report'])) & ~tr['StudyInstanceUID'].isin(val_ids)\ntr = tr[~leak].reset_index(drop=True)\ntr = tr[tr['StudyInstanceUID'].isin(soft.index) | tr['StudyInstanceUID'].isin(val_ids)].reset_index(drop=True)\nfor c in LABELS:\n    tr[c] = tr['StudyInstanceUID'].map(soft[c]).astype(float)\ntr = tr.dropna(subset=LABELS).reset_index(drop=True)\n\ndup = tr['Report'].duplicated(keep=False)\npool = tr[~tr['StudyInstanceUID'].isin(val_ids) & ~dup]\nhold_ids = set(pool.sample(250, random_state=7)['StudyInstanceUID'])   # same seed as v4 -> comparable\n\nslot_map = {}\nfor study, g in ser.groupby('StudyInstanceUID'):\n    row = {}\n    for name, cond in SLOTS:\n        m = g[cond(g)]\n        if len(m): row[name] = m.iloc[0]['SeriesInstanceUID']\n    slot_map[study] = row\ntr = tr[tr['StudyInstanceUID'].isin(slot_map)].reset_index(drop=True)\nlog(f'{len(tr)} studies | {len(val_ids)} annotated | {len(hold_ids)} holdout')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def get_laterality(study):\n    for name in slot_map[study]:\n        for f in glob.glob(os.path.join(PATH, 'train_series', study, slot_map[study][name], '*.dcm'))[:2]:\n            try:\n                ds = pydicom.dcmread(f, stop_before_pixels=True)\n                lat = str(getattr(ds, 'Laterality', '') or getattr(ds, 'ImageLaterality', '')).upper()\n                if lat in ('L', 'R'): return lat\n            except Exception:\n                pass\n    return 'L'\nwith ThreadPoolExecutor(16) as ex:\n    lats = list(ex.map(get_laterality, tr['StudyInstanceUID']))\nlat_map = dict(zip(tr['StudyInstanceUID'], lats))\nlog(f\"laterality: {pd.Series(lats).value_counts().to_dict()}\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"studies = list(tr['StudyInstanceUID'])\nsidx = {s: i for i, s in enumerate(studies)}\nCACHE = np.zeros((len(studies), len(SLOTS), N_SLICES, CACHE_SIZE, CACHE_SIZE), dtype=np.uint8)\nMASK = np.zeros((len(studies), len(SLOTS)), dtype=np.float32)\n\ndef decode(study):\n    i = sidx[study]; lat = lat_map[study]\n    for j, (name, _) in enumerate(SLOTS):\n        if name not in slot_map[study]: continue\n        sdir = os.path.join(PATH, 'train_series', study, slot_map[study][name])\n        slices = []\n        for f in glob.glob(os.path.join(sdir, '*.dcm')):\n            try:\n                ds = pydicom.dcmread(f)\n                slices.append((int(getattr(ds, 'InstanceNumber', 0)), ds.pixel_array.astype(np.float32)))\n            except Exception:\n                continue\n        if not slices: continue\n        slices.sort(key=lambda x: x[0])\n        if lat == 'R' and name == 'sagfs': slices = slices[::-1]\n        take = np.linspace(0, len(slices)-1, N_SLICES).round().astype(int)\n        for k, t in enumerate(take):\n            img = slices[t][1]\n            lo, hi = np.percentile(img, [1, 99])\n            img = np.clip((img - lo) / max(hi - lo, 1e-6), 0, 1)\n            img = cv2.resize(img, (CACHE_SIZE, CACHE_SIZE), interpolation=cv2.INTER_AREA)\n            if lat == 'R' and name in ('cor', 'ax'): img = img[:, ::-1]\n            CACHE[i, j, k] = (img * 255).astype(np.uint8)\n        MASK[i, j] = 1.0\n\nwith ThreadPoolExecutor(12) as ex:\n    for n, _ in enumerate(ex.map(decode, studies)):\n        if (n+1) % 1000 == 0: log(f'cached {n+1}/{len(studies)}')\nlog(f'cache {CACHE.nbytes/1e9:.1f} GB')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"Y = tr[LABELS].values.astype(np.float32)\nis_val = tr['StudyInstanceUID'].isin(val_ids).values\nis_hold = tr['StudyInstanceUID'].isin(hold_ids).values\nva_order = [sidx[s] for s in sorted(val_ids) if s in sidx]\nho_order = np.where(is_hold)[0]\n\nclass DS(Dataset):\n    def __init__(self, idxs, train): self.idxs, self.train = idxs, train\n    def __len__(self): return len(self.idxs)\n    def __getitem__(self, k):\n        i = self.idxs[k]\n        v = CACHE[i].astype(np.float32) / 255.\n        if self.train:\n            ox, oy = np.random.randint(0, CACHE_SIZE-SIZE+1, 2)\n            v = v[:, :, oy:oy+SIZE, ox:ox+SIZE] * np.random.uniform(0.85, 1.15)\n            v = v + np.random.uniform(-0.05, 0.05)\n        else:\n            o = (CACHE_SIZE-SIZE)//2\n            v = v[:, :, o:o+SIZE, o:o+SIZE]\n        v = ((np.clip(v, 0, 1) - 0.45) / 0.225).astype(np.float32)\n        return torch.from_numpy(np.ascontiguousarray(v)), torch.from_numpy(MASK[i]), torch.from_numpy(Y[i])\n\ntr_idx = np.where(~is_val & ~is_hold)[0]\ntl = DataLoader(DS(tr_idx, True), batch_size=BATCH, shuffle=True, num_workers=3, pin_memory=True, drop_last=True)\nvl = DataLoader(DS(va_order, False), batch_size=4, shuffle=False, num_workers=2)\nhl = DataLoader(DS(ho_order, False), batch_size=4, shuffle=False, num_workers=2)\nlog(f'train {len(tr_idx)} | val {len(va_order)} | holdout {len(ho_order)}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def radimagenet_backbone():\n    \"\"\"torchvision ResNet-50 trunk loaded with RadImageNet weights (keys are backbone.0..backbone.7).\"\"\"\n    r = torchvision.models.resnet50(weights=None)\n    trunk = nn.Sequential(*list(r.children())[:-2])          # conv1..layer4 -> (B, 2048, H/32, W/32)\n    sd = torch.load(RADW, map_location='cpu')\n    sd = {k[len('backbone.'):]: v for k, v in sd.items() if k.startswith('backbone.')}\n    missing, unexpected = trunk.load_state_dict(sd, strict=False)\n    print(f'RadImageNet load -> missing {len(missing)}, unexpected {len(unexpected)}')\n    assert len(missing) == 0, missing[:5]\n    return trunk\n\nclass SliceAttnPool(nn.Module):\n    def __init__(self, dim, hidden=256):\n        super().__init__()\n        self.V = nn.Linear(dim, hidden); self.U = nn.Linear(dim, hidden); self.w = nn.Linear(hidden, 1)\n    def forward(self, f):\n        a = self.w(torch.tanh(self.V(f)) * torch.sigmoid(self.U(f)))\n        return (torch.softmax(a, dim=2) * f).sum(2)\n\nclass RadPlaneNet(nn.Module):\n    def __init__(self, n_slots=3, n_out=12):\n        super().__init__()\n        self.bb = radimagenet_backbone()\n        D = 2048\n        self.pool = SliceAttnPool(D)\n        self.slot_emb = nn.Parameter(torch.zeros(n_slots, D))\n        self.drop = nn.Dropout(0.3)\n        self.expert = nn.Linear(D, n_out); self.gate = nn.Linear(D, n_out)\n    def forward(self, x, mask):\n        B, S, C, H, W = x.shape\n        f = self.bb(x.reshape(B*S*C, 1, H, W).expand(-1, 3, -1, -1))\n        f = torch.nn.functional.adaptive_avg_pool2d(f, 1).flatten(1)\n        f = self.pool(f.reshape(B, S, C, -1)) + self.slot_emb\n        f = self.drop(f)\n        scores = self.gate(f).masked_fill(mask.unsqueeze(-1) == 0, -1e4)\n        return (torch.softmax(scores, 1) * self.expert(f)).sum(1)\n\nmodel = RadPlaneNet().to(DEVICE)\nbb_p = [p for n, p in model.named_parameters() if n.startswith('bb.')]\nhd_p = [p for n, p in model.named_parameters() if not n.startswith('bb.')]\nopt = torch.optim.AdamW([{'params': bb_p, 'lr': 6e-5}, {'params': hd_p, 'lr': 4e-4}], weight_decay=0.05)\nspe = len(tl) // ACCUM; total = EPOCHS * spe; warm = spe\nsched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: (s+1)/warm if s < warm else 0.5*(1+np.cos(np.pi*(s-warm)/max(1, total-warm))))\ncrit = nn.BCEWithLogitsLoss()\nscaler = torch.cuda.amp.GradScaler()\nema = copy.deepcopy(model).eval()\nfor p in ema.parameters(): p.requires_grad_(False)\ndef ema_update(d=0.998):\n    with torch.no_grad():\n        for pe, pm in zip(ema.parameters(), model.parameters()): pe.mul_(d).add_(pm.detach(), alpha=1-d)\n        for be, bm in zip(ema.buffers(), model.buffers()): be.copy_(bm)\nlog('model ready')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def predict(net, loader):\n    net.eval(); P = []\n    with torch.no_grad():\n        for x, m, y in loader:\n            with torch.cuda.amp.autocast():\n                P.append(torch.sigmoid(net(x.to(DEVICE), m.to(DEVICE))).float().cpu().numpy())\n    return np.concatenate(P)\n\ndef macro(P, YT):\n    a = {}\n    for j, c in enumerate(LABELS):\n        yj = (YT[:, j] >= 0.5).astype(int)\n        if 0 < yj.sum() < len(yj): a[c] = roc_auc_score(yj, P[:, j])\n    return float(np.mean(list(a.values()))), a\n\nyt_val = y_true.loc[[studies[i] for i in va_order]].values\nyt_hold = Y[ho_order]\nbest_h = 0.0\nfor ep in range(EPOCHS):\n    model.train(); tot = 0.0; t0 = time.time(); opt.zero_grad()\n    for k, (x, m, y) in enumerate(tl):\n        with torch.cuda.amp.autocast():\n            loss = crit(model(x.to(DEVICE, non_blocking=True), m.to(DEVICE)), y.to(DEVICE)) / ACCUM\n        scaler.scale(loss).backward(); tot += loss.item()*ACCUM*len(x)\n        if (k+1) % ACCUM == 0:\n            scaler.step(opt); scaler.update(); opt.zero_grad(); sched.step(); ema_update()\n    mh, _ = macro(predict(ema, hl), yt_hold)\n    mv, av = macro(predict(ema, vl), yt_val)\n    flag = ''\n    if mh > best_h:                      # SELECT ON HOLDOUT, not the 58\n        best_h = mh; torch.save(ema.state_dict(), 'model_v5_best.pt'); flag = '  <-- best saved (holdout)'\n    log(f'epoch {ep+1}: loss {tot/len(tr_idx):.4f} | holdout {mh:.4f} | annotated-val {mv:.4f} ({(time.time()-t0)/60:.1f} min){flag}')\n    for c, a in sorted(av.items(), key=lambda kv: kv[1]): print(f'   {c}: {a:.3f}')\ntorch.save(ema.state_dict(), 'model_v5_last.pt')\nlog(f'BEST holdout {best_h:.4f}')"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"}},"nbformat":4,"nbformat_minor":5}