{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee Swin-Tiny dense multi-span\n\nFrozen offline Swin-Tiny with dense geometry-ordered slice coverage.\nThe dense arm is combined with the established single and three-slice arms.\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"# All dependencies and weights are attached as Kaggle inputs.\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from pathlib import Path\nimport random, warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\nimport torch\nfrom torch import nn\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ndevice = 'cpu'\nlabels = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nroot = next(p for p in Path('/kaggle/input').glob('**/train.csv') if (p.parent / 'sample_submission.csv').exists())\ntrain = pd.read_csv(root); test = pd.read_csv(root.parent / 'test.csv')\ntrain_series = pd.read_csv(root.parent / 'train_series.csv'); test_series = pd.read_csv(root.parent / 'test_series.csv')\ngold = train.dropna(subset=labels).reset_index(drop=True)\nprint('gold studies:', len(gold), 'test studies:', len(test))\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"def choose_series(df):\n    x = df.copy()\n    x['_score'] = 2 * pd.to_numeric(x.get('Fluid_Sensitive', 0), errors='coerce').fillna(0) + pd.to_numeric(x.get('Fat_Suppression', 0), errors='coerce').fillna(0)\n    plane = x.get('Anatomical_Plane', '').fillna('').astype(str).str.lower()\n    x.loc[plane.str.contains('sagittal'), '_score'] += 2\n    x.loc[plane.str.contains('coronal'), '_score'] += 1\n    selected = {}\n    for study, group in x.groupby('StudyInstanceUID'):\n        chosen = []\n        for wanted in ['sagittal', 'coronal', 'axial']:\n            match = group[group['Anatomical_Plane'].fillna('').astype(str).str.lower().eq(wanted)]\n            if len(match): chosen.append(match.sort_values('_score', ascending=False).iloc[0]['SeriesInstanceUID'])\n        if len(chosen) < 3:\n            for sid in group.sort_values('_score', ascending=False)['SeriesInstanceUID']:\n                if sid not in chosen: chosen.append(sid)\n                if len(chosen) == 3: break\n        selected[study] = chosen[:3]\n    return selected\n\ndef ordered_dicoms(folder):\n    rows = []\n    for path in Path(folder).glob('*.dcm'):\n        try:\n            ds = pydicom.dcmread(str(path), stop_before_pixels=True, force=True)\n            pos = np.asarray(getattr(ds, 'ImagePositionPatient', []), dtype='float32')\n            ori = np.asarray(getattr(ds, 'ImageOrientationPatient', []), dtype='float32')\n            if pos.size == 3 and ori.size == 6:\n                coord = float(np.dot(pos, np.cross(ori[:3], ori[3:])))\n            else: coord = float(getattr(ds, 'InstanceNumber', 0))\n        except Exception: coord = 0.0\n        rows.append((coord, path))\n    return [path for _, path in sorted(rows, key=lambda item: item[0])]\n\ndef read_slice(path, size=224):\n    ds = pydicom.dcmread(str(path), force=True)\n    arr = ds.pixel_array.astype('float32')\n    arr = arr * float(getattr(ds, 'RescaleSlope', 1.0)) + float(getattr(ds, 'RescaleIntercept', 0.0))\n    lo, hi = np.percentile(arr, [1, 99]) if arr.max() > arr.min() else (arr.min(), arr.min() + 1)\n    arr = np.clip((arr - lo) / (hi - lo), 0, 1)\n    return np.asarray(Image.fromarray((arr * 255).astype('uint8')).resize((size, size)), dtype='uint8')\n\ndef study_image(study, series_map, image_root, offset=0):\n    images = []\n    for sid in series_map.get(study, []):\n        paths = ordered_dicoms(image_root / str(study) / str(sid))\n        if paths:\n            center = len(paths) // 2\n            idx = min(max(center + offset, 0), len(paths) - 1)\n            images.append(read_slice(paths[idx]))\n    while len(images) < 3: images.append(np.zeros((224, 224), dtype='uint8'))\n    return np.stack(images[:3], axis=-1)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from transformers import AutoModel\nmodel_dir = next(p for p in Path('/kaggle/input').glob('**/rainduck32/swin-tiny-patch4-window7-224/**') if p.is_dir() and (p / 'config.json').exists())\nmodel = AutoModel.from_pretrained(str(model_dir), local_files_only=True).to(device).eval()\nseries_tr = choose_series(train_series); series_te = choose_series(test_series)\ntrain_root = root.parent / 'train_series'; test_root = root.parent / 'test_series'\n\ndef encode(frame, series_map, image_root, offsets):\n    feats = []\n    with torch.no_grad():\n        for j, study in enumerate(frame['StudyInstanceUID']):\n            views = []\n            for offset in offsets:\n                image = Image.fromarray(study_image(study, series_map, image_root, offset))\n                x = torch.tensor(np.asarray(image), dtype=torch.float32).permute(2, 0, 1) / 255.0\n                x = (x - 0.5) / 0.5\n                out = model(pixel_values=x.unsqueeze(0))\n                feat = getattr(out, 'pooler_output', None)\n                if feat is None: feat = out.last_hidden_state.mean(dim=1)\n                views.append(feat.float().cpu().numpy()[0])\n            feats.append(np.mean(views, axis=0))\n            if (j + 1) % 20 == 0: print('encoded', j + 1, '/', len(frame), 'offsets', len(offsets))\n    return np.asarray(feats, dtype='float32')\n\nX_single = encode(gold, series_tr, train_root, (0,))\nXt_single = encode(test, series_te, test_root, (0,))\nX_multi = encode(gold, series_tr, train_root, (-2, 0, 2))\nXt_multi = encode(test, series_te, test_root, (-2, 0, 2))\nDENSE_OFFSETS = (-12, -9, -6, -3, 0, 3, 6, 9, 12)\nX_dense = encode(gold, series_tr, train_root, DENSE_OFFSETS)\nXt_dense = encode(test, series_te, test_root, DENSE_OFFSETS)\nprint(X_single.shape, X_multi.shape, X_dense.shape, Xt_dense.shape)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from sklearn.metrics import roc_auc_score\nrng = np.random.default_rng(SEED)\nperm = rng.permutation(len(gold)); cut = max(1, int(len(gold) * 0.8))\ntr_idx, va_idx = perm[:cut], perm[cut:]\nyy = torch.tensor(gold[labels].to_numpy(dtype='float32'))\n\ndef macro_auc(y_true, pred):\n    vals = [roc_auc_score(y_true[:,j], pred[:,j]) for j in range(y_true.shape[1]) if len(np.unique(y_true[:,j])) == 2]\n    return float(np.mean(vals)) if vals else 0.5\n\ndef select_probe(X):\n    xx = torch.tensor(X); probe = nn.Linear(X.shape[1], len(labels))\n    opt = torch.optim.AdamW(probe.parameters(), lr=5e-4, weight_decay=5e-2)\n    best_loss, best_epoch, best_state, patience = float('inf'), 1, None, 0\n    for epoch in range(180):\n        loss = nn.functional.binary_cross_entropy_with_logits(probe(xx[tr_idx]), yy[tr_idx])\n        opt.zero_grad(); loss.backward(); opt.step()\n        with torch.no_grad(): val_loss = nn.functional.binary_cross_entropy_with_logits(probe(xx[va_idx]), yy[va_idx]).item()\n        if val_loss < best_loss - 1e-4:\n            best_loss, best_epoch, patience = val_loss, epoch + 1, 0\n            best_state = {k:v.detach().cpu().clone() for k,v in probe.state_dict().items()}\n        else: patience += 1\n        if patience >= 35: break\n    return best_epoch, best_state\n\ndef fit_full(X, epochs):\n    xx = torch.tensor(X); head = nn.Linear(X.shape[1], len(labels))\n    opt = torch.optim.AdamW(head.parameters(), lr=5e-4, weight_decay=5e-2)\n    for _ in range(epochs):\n        loss = nn.functional.binary_cross_entropy_with_logits(head(xx), yy)\n        opt.zero_grad(); loss.backward(); opt.step()\n    return head\n\ndef pred(X, state):\n    h = nn.Linear(X.shape[1], len(labels)); h.load_state_dict(state); h.eval()\n    with torch.no_grad(): return torch.sigmoid(h(torch.tensor(X[va_idx]))).numpy()\n\narms = {'single': X_single, 'multi': X_multi, 'dense': X_dense}\nstates = {}; epochs = {}; ranks = {}\nfor name, X_arm in arms.items():\n    epochs[name], states[name] = select_probe(X_arm)\n    ranks[name] = pd.DataFrame(pred(X_arm, states[name])).rank(pct=True).to_numpy()\ny_val = gold[labels].to_numpy(dtype='int32')[va_idx]\ncandidates = [(a,b,c) for a in (0,0.25,0.5,0.75,1) for b in (0,0.25,0.5,0.75,1) for c in (0,0.25,0.5,0.75,1) if abs(a+b+c-1) < 1e-6]\nscores = {w: macro_auc(y_val, w[0]*ranks['single'] + w[1]*ranks['multi'] + w[2]*ranks['dense']) for w in candidates}\nbest = max(scores, key=scores.get)\nprint('dense arm validation losses/epochs:', epochs)\nprint('best dense ensemble:', best, 'validation AUC:', round(scores[best], 4))\nheads = {name: fit_full(X_arm, epochs[name]) for name, X_arm in arms.items()}\ntest_arms = {'single': Xt_single, 'multi': Xt_multi, 'dense': Xt_dense}\ntest_rank = {}\nfor name, X_arm in test_arms.items():\n    with torch.no_grad(): out = torch.sigmoid(heads[name](torch.tensor(X_arm))).numpy()\n    test_rank[name] = pd.DataFrame(out).rank(pct=True).to_numpy()\npred_final = sum(w * test_rank[name] for w, name in zip(best, ('single','multi','dense')))\nsubmission = pd.read_csv(root.parent / 'sample_submission.csv'); submission[labels] = pred_final; submission.to_csv('/kaggle/working/submission.csv', index=False)\ndisplay(submission.head())\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"}},"nbformat":4,"nbformat_minor":5}