{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee Swin-Tiny weak-supervised\n\nFrozen offline Swin-Tiny encoder trained on fold-safe report targets.\nThe 58 complete expert-labeled studies override weak targets with higher weight.\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, hashlib\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_mask = train[labels].notna().all(axis=1)\ngold = train.loc[gold_mask].reset_index(drop=True)\nprint('all studies:', len(train), '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=(0,)):\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) % 200 == 0: print('encoded', j + 1, '/', len(frame))\n    return np.asarray(feats, dtype='float32')\n\nX = encode(train, series_tr, train_root)\nXt = encode(test, series_te, test_root)\nprint(X.shape, Xt.shape)\n"},{"cell_type":"code","metadata":{},"execution_count":null,"outputs":[],"source":"from sklearn.metrics import roc_auc_score\n\nPATTERNS = {\n    'ACL': r'\\bacl\\b|lca|ligament.*cruzad.*anter|ligament croise|vkb',\n    'MCL': r'\\bmcl\\b|collateral.*medial|colateral medial|innenband',\n    'Medial Meniscus': r'menisc.*(medial|intern)|(medial|intern).*menisc',\n    'Lateral Meniscus': r'menisc.*(lateral|extern)|(lateral|extern).*menisc',\n    'Medial OA': r'(arthros|osteoarth|artros).{0,20}(medial|intern)|(medial|intern).{0,20}(arthros|osteoarth|artros)',\n    'Lateral OA': r'(arthros|osteoarth|artros).{0,20}(lateral|extern)|(lateral|extern).{0,20}(arthros|osteoarth|artros)',\n    'PF OA': r'(patellofemoral|femoropatellar|rotulofemoral).{0,20}(arthros|osteoarth|artros)',\n    'Effusion': r'effusion|joint fluid|gelenkerguss|hydarthr|epanchement',\n    'Synovitis': r'synov|plica|hoff', 'Baker\\'s': r'baker|popliteal cyst|kyste poplite',\n    'Contusion': r'contus|bone bruise|bone marrow edema|knochenmark',\n    'Fracture': r'fractur|fraktur|frattur|insufficiency fracture',\n}\nNEG = r'(?:no|not|without|negative for|nicht|keine?|kein|aucun|sans|sin|non|absence de)\\W{0,45}$'\ndef weak_states(text):\n    import re, unicodedata\n    text = unicodedata.normalize('NFKD', str(text)).encode('ascii', 'ignore').decode('ascii').lower()\n    out = []\n    for label in labels:\n        state = 0\n        for m in re.finditer(PATTERNS[label], text):\n            state = -1 if re.search(NEG, text[max(0, m.start()-55):m.start()]) else 1\n            if state == 1: break\n        out.append(state)\n    return out\n\nstates = np.asarray([weak_states(x) for x in train['Report'].fillna('')], dtype='int8')\nY = np.where(states == 1, 0.85, np.where(states == -1, 0.15, 0.5)).astype('float32')\nW = np.where(states == 0, 0.15, 1.0).astype('float32')\nY[gold_mask.to_numpy()] = gold[labels].to_numpy(dtype='float32')\nW[gold_mask.to_numpy()] = 8.0\ngroups = np.asarray([int(hashlib.md5(str(x).encode()).hexdigest()[:8], 16) % 5 for x in train['Report'].fillna('')])\ntr_idx = np.flatnonzero(groups != 0); va_idx = np.flatnonzero(groups == 0)\nva_gold_idx = np.flatnonzero((groups == 0) & gold_mask.to_numpy())\nif len(va_gold_idx) == 0: raise RuntimeError('validation fold has no complete gold studies')\nyy = torch.tensor(Y); ww = torch.tensor(W)\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 train_head(epochs, indices, validate=False):\n    head = nn.Linear(X.shape[1], len(labels))\n    opt = torch.optim.AdamW(head.parameters(), lr=5e-4, weight_decay=5e-2)\n    best, state, patience = -1.0, None, 0\n    for ep in range(epochs):\n        logits = head(torch.tensor(X[indices]))\n        loss = (nn.functional.binary_cross_entropy_with_logits(logits, yy[indices], reduction='none') * ww[indices]).mean()\n        opt.zero_grad(); loss.backward(); opt.step()\n        if validate:\n            with torch.no_grad(): val = torch.sigmoid(head(torch.tensor(X[va_gold_idx]))).numpy()\n            score = macro_auc(train.loc[va_gold_idx, labels].to_numpy(dtype='float32'), val)\n            if score > best + 1e-4:\n                best, state, patience = score, {k:v.detach().cpu().clone() for k,v in head.state_dict().items()}, 0\n            else: patience += 1\n            if patience >= 35: break\n    return ep + 1, best, state, head\n\nepochs, val_score, state, _ = train_head(180, tr_idx, validate=True)\nhead = nn.Linear(X.shape[1], len(labels)); head.load_state_dict(state)\n_, _, _, head = train_head(max(epochs, 1), np.arange(len(train)), validate=False)\nwith torch.no_grad(): pred = torch.sigmoid(head(torch.tensor(Xt))).numpy()\nsubmission = pd.read_csv(root.parent / 'sample_submission.csv')\nsubmission[labels] = pd.DataFrame(pred).rank(pct=True).to_numpy()\nsubmission.to_csv('/kaggle/working/submission.csv', index=False)\nprint('weak validation AUC:', round(val_score, 4), 'epochs:', epochs)\ndisplay(submission.head())\n"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"}},"nbformat":4,"nbformat_minor":5}