{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee Swin-Tiny balanced head\n\nFrozen offline Swin-Tiny encoder with a class-balanced multilabel head.\nThe balance weights are estimated from the training portion only.\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\n\nmodel_dir = next(\n    p for p in Path('/kaggle/input').glob('**/rainduck32/swin-tiny-patch4-window7-224/**')\n    if p.is_dir() and (p / 'config.json').exists()\n)\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            study_feats = []\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                study_feats.append(feat.float().cpu().numpy()[0])\n            feats.append(np.mean(study_feats, axis=0))\n            if (j + 1) % 20 == 0: print('encoded', j + 1, '/', len(frame), 'offsets', 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))\nprint(X_single.shape, X_multi.shape, Xt_single.shape, Xt_multi.shape)\n\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 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\ndef class_weights(indices):\n    y = gold[labels].to_numpy(dtype='float32')[indices]\n    pos = y.sum(axis=0)\n    neg = len(indices) - pos\n    ratio = neg / np.maximum(pos, 1.0)\n    return torch.tensor(np.clip(ratio, 0.5, 3.0), dtype=torch.float32)\n\ndef select_probe(X):\n    xx = torch.tensor(X)\n    probe = nn.Linear(X.shape[1], len(labels))\n    pos_weight = class_weights(tr_idx)\n    opt = torch.optim.AdamW(probe.parameters(), lr=5e-4, weight_decay=5e-2)\n    best_score, best_epoch, best_state, patience = -1.0, 1, None, 0\n    for epoch in range(220):\n        loss = nn.functional.binary_cross_entropy_with_logits(\n            probe(xx[tr_idx]), yy[tr_idx], pos_weight=pos_weight)\n        opt.zero_grad(); loss.backward(); opt.step()\n        with torch.no_grad():\n            val = torch.sigmoid(probe(xx[va_idx])).numpy()\n        score = macro_auc(gold[labels].to_numpy(dtype='int32')[va_idx], val)\n        if score > best_score + 1e-4:\n            best_score, best_epoch, patience = score, 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 >= 40: break\n    return best_epoch, best_score, best_state\n\ndef predict_state(X, state, rows):\n    head = nn.Linear(X.shape[1], len(labels))\n    head.load_state_dict(state); head.eval()\n    with torch.no_grad():\n        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    weight = class_weights(np.arange(len(gold)))\n    for _ in range(epochs):\n        loss = nn.functional.binary_cross_entropy_with_logits(\n            head(xx), yy, pos_weight=weight)\n        opt.zero_grad(); loss.backward(); opt.step()\n    return head\n\nep_single, score_single, state_single = select_probe(X_single)\nep_multi, score_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()\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('balanced single validation AUC:', round(score_single, 4), 'epochs:', ep_single)\nprint('balanced multi validation AUC:', round(score_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}