{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee Swin-Tiny six-view fusion\n\nFrozen offline Swin-Tiny with explicit fusion of three planes and two contrast types.\nEach study contributes six separate view embeddings before multilabel prediction.\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['_fluid'] = pd.to_numeric(x['Fluid_Sensitive'], errors='coerce').fillna(-1).astype(int)\n    x['_plane'] = x['Anatomical_Plane'].fillna('').astype(str).str.lower()\n    selected = {}\n    fallback_count = 0\n    for study, group in x.groupby('StudyInstanceUID'):\n        chosen = []\n        for fluid in [0, 1]:\n            for wanted in ['sagittal', 'coronal', 'axial']:\n                match = group[(group['_fluid'] == fluid) & (group['_plane'] == wanted)]\n                if len(match):\n                    chosen.append(match.sort_values('SeriesInstanceUID').iloc[0]['SeriesInstanceUID'])\n                else:\n                    same_fluid = group[group['_fluid'] == fluid]\n                    pool = same_fluid if len(same_fluid) else group\n                    if not len(pool): raise RuntimeError(f'study {study} has no series')\n                    chosen.append(pool.sort_values('SeriesInstanceUID').iloc[0]['SeriesInstanceUID'])\n                    fallback_count += 1\n        selected[study] = chosen\n    print('explicit plane-contrast fallbacks:', fallback_count, '/', len(selected) * 6)\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            coord = float(np.dot(pos, np.cross(ori[:3], ori[3:]))) if pos.size == 3 and ori.size == 6 else 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, view, offset=0):\n    sid = series_map[study][view]\n    paths = ordered_dicoms(image_root / str(study) / str(sid))\n    if not paths: raise RuntimeError(f'empty series {sid}')\n    center = len(paths) // 2\n    idx = min(max(center + offset, 0), len(paths) - 1)\n    return read_slice(paths[idx])\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'\nprint('six-view coverage:', len(series_tr), len(series_te), 'views per study:', sorted(set(map(len, series_tr.values()))))\n\ndef encode(frame, series_map, image_root):\n    feats = []\n    with torch.no_grad():\n        for j, study in enumerate(frame['StudyInstanceUID']):\n            study_feats = []\n            for view in range(6):\n                gray = study_image(study, series_map, image_root, view)\n                image = Image.fromarray(np.repeat(gray[..., None], 3, axis=2))\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                study_feats.append(feat.float().cpu().numpy()[0])\n            feats.append(np.concatenate(study_feats, axis=0))\n            if (j + 1) % 200 == 0: print('encoded', j + 1, '/', len(frame))\n    return np.asarray(feats, dtype='float32')\n\nX = encode(gold, series_tr, train_root)\nXt = encode(test, series_te, test_root)\nX_single, Xt_single, X_multi, Xt_multi = X, Xt, X, Xt\nprint(X.shape, Xt.shape)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from sklearn.metrics import roc_auc_score\n\nrng = np.random.default_rng(SEED)\nperm = rng.permutation(len(gold))\ncut = 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 select_probe(X):\n    xx = torch.tensor(X)\n    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(\n            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:\n            patience += 1\n        if patience >= 35: break\n    return best_epoch, best_loss, best_state\n\ndef predict_state(X, state, rows):\n    head = nn.Linear(X.shape[1], len(labels))\n    head.load_state_dict(state)\n    head.eval()\n    with torch.no_grad(): return torch.sigmoid(head(torch.tensor(X[rows]))).numpy()\n\ndef fit_full(X, epochs):\n    xx = torch.tensor(X)\n    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\nep_single, loss_single, state_single = select_probe(X_single)\nep_multi, loss_multi, state_multi = select_probe(X_multi)\nval_single = predict_state(X_single, state_single, va_idx)\nval_multi = predict_state(X_multi, state_multi, va_idx)\ny_val = gold[labels].to_numpy(dtype='int32')[va_idx]\nrank_single = pd.DataFrame(val_single).rank(pct=True).to_numpy()\nrank_multi = pd.DataFrame(val_multi).rank(pct=True).to_numpy()\n\ndef macro_auc(y_true, pred):\n    scores = []\n    for j in range(y_true.shape[1]):\n        if len(np.unique(y_true[:, j])) == 2:\n            scores.append(roc_auc_score(y_true[:, j], pred[:, j]))\n    return float(np.mean(scores)) if scores else 0.5\n\nweights = [0.0, 0.25, 0.5, 0.75, 1.0]\nscores = {w: macro_auc(y_val, w * rank_single + (1 - w) * rank_multi) for w in weights}\nbest_weight = max(scores, key=scores.get)\nprint('single validation loss:', round(loss_single, 4), 'epochs:', ep_single)\nprint('multi validation loss:', round(loss_multi, 4), 'epochs:', ep_multi)\nprint('rank ensemble validation scores:', {str(k): round(v, 4) for k, v in scores.items()})\nprint('selected single weight:', best_weight)\n\nhead_single = fit_full(X_single, ep_single)\nhead_multi = fit_full(X_multi, ep_multi)\nwith torch.no_grad():\n    pred_single = torch.sigmoid(head_single(torch.tensor(Xt_single))).numpy()\n    pred_multi = torch.sigmoid(head_multi(torch.tensor(Xt_multi))).numpy()\nrank_single = pd.DataFrame(pred_single).rank(pct=True).to_numpy()\nrank_multi = pd.DataFrame(pred_multi).rank(pct=True).to_numpy()\npred_ensemble = best_weight * rank_single + (1 - best_weight) * rank_multi\nsubmission = pd.read_csv(root.parent / 'sample_submission.csv')\nsubmission[labels] = pred_ensemble\nsubmission.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}