{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":154281},{"sourceType":"datasetVersion","sourceId":18673450},{"sourceType":"modelInstanceVersion","sourceId":4533}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"papermill":{"default_parameters":{},"duration":32319.421319,"end_time":"2026-09-19T15:05:48.538645+00:00","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-09-19T06:07:09.117326+00:00","version":"2.7.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"d57a09af-1a38-401b-88c2-8d78c2b38a8b","cell_type":"markdown","source":"# RSNA Knee — training notebook, v2\n\nSame pipeline as the original, plus: model presets, the embedding size read from the backbone, a saved fold file so several models validate on the same studies, saved soft targets, an optional label candidate (average of lexicon and label table), an optional non-optimistic Phase A selection, and a per-finding AUC report with a bootstrap interval. Defaults reproduce the original behaviour except `LABEL_BLEND`, which is only adopted if it beats the lexicon by more than 0.01 on the gold studies.\n\n**Not run end to end by the author of these edits** (no GPU or competition data was available); each edit was syntax-checked only.","metadata":{}},{"id":"dee11b61","cell_type":"markdown","source":"# RSNA 2026 Knee Abnormality Detection\n\n\n\n","metadata":{"papermill":{"duration":0.005194,"end_time":"2026-09-19T06:07:11.716746+00:00","exception":false,"start_time":"2026-09-19T06:07:11.711552+00:00","status":"completed"},"tags":[]}},{"id":"b492cd9b","cell_type":"code","source":"DATA_DIR  = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\nMODEL_DIR = \"/kaggle/input/models/metaresearch/dinov2/pytorch/small/1\"\nWORK      = \"/kaggle/working\"\nCACHE_DIR = \"/kaggle/working\"   # point at a mounted previous output to skip 72 min\nCKPT_DIR  = \"/kaggle/working\"\nLABEL_DIR   = \"/kaggle/input/datasets/pilkwang/rsna-knee-llm-labels\"\nLABEL_TABLE = f\"{LABEL_DIR}/report_labels_v2.csv\"   # falls back to scanning LABEL_DIR\nREQUIRE_LABELS = True    # stop if LABEL_DIR is set but no usable table is found there   # infer kernel: the mounted training output\nSTAGE     = \"train\"\n\nTIMM_NAME = \"vit_small_patch14_dinov2.lvd142m\"\nEMB_DIM   = None         # set from the backbone in cell 8 (3 x embed dim: CLS + mean-patch + top-k/8 patch mean)\nIMG_SIZE  = 336\nCROP_MM   = 130.0\n\nSLOT_DEFS  = [(\"axial\", 1), (\"sagittal\", 0), (\"coronal\", 1), (\"sagittal\", 1)]\nK_PER_SLOT = [8, 12, 10, 16]\nFP_FIELDS  = [\"Manufacturer\", \"ManufacturerModelName\", \"SoftwareVersions\", \"ReceiveCoilName\"]\n\nN_FOLDS, FOLDS = 5, [0, 1, 2, 3, 4]\nHEAD_EPOCHS    = 120\nPHASE_B_EPOCHS = 5\nBATCH_STUDIES  = 6\nACCUM          = 3\nEXTRACT_BATCH  = 8\nNUM_WORKERS    = 4\nHDR_THREADS    = 8       # per DataLoader worker, for the slice-ordering header pass\nLR_HEAD, LR_BB = 1e-3, 5e-5\nWD, UNFREEZE   = 0.01, 8\nEMA_DECAY      = 0.99\nGOLD_W         = 3.0\nHEAD_HIDDEN    = 256\nHEAD_DROP      = 0.2\nPRIOR_W        = 1.0     # anatomical bias on the per-label attention logits\nMAX_MIX        = 0.30\nLAB_MIX        = 0.15\nTXT_W          = 0.10\nRUN_PHASE_B    = True\nSAVE_JPEG      = True\nKEEP_JPEG      = False   # True keeps the 5.4 GB cache so a rerun can resume mid-training\nJPEG_Q         = 92\nTIME_BUDGET    = 10.5 * 3600\nFOLD_SECONDS   = 105 * 60\nSEED           = 42\n\nTTA      = 3\nEFF_MODE = False\nEFF_IMG, EFF_K = 224, [5, 7, 6, 9]\n\n# ---------------------------------------------------------------- v2 additions\n# PRESET picks which model this run trains. \"small336\" is the original notebook, unchanged.\n#   small280_seed2 : same DINOv2-small, 280 px, other seed  -> cheap companion model (~0.7x the compute)\n#   base_frozen    : DINOv2-base, frozen features + head only (no Phase B) -> tests a bigger backbone quickly\n# Everything else stays as configured above. Check the base model folder name under /kaggle/input/models.\nPRESET = \"small336\"\nPRESETS = {\n    \"small336\": {},\n    \"small280_seed2\": dict(IMG_SIZE=280, SEED=7),\n    \"base_frozen\": dict(\n        MODEL_DIR=\"/kaggle/input/models/metaresearch/dinov2/pytorch/base/1\",\n        TIMM_NAME=\"vit_base_patch14_dinov2.lvd142m\",\n        RUN_PHASE_B=False, SAVE_JPEG=False, SEED=7),\n}\nglobals().update(PRESETS[PRESET])\n\nFOLD_SRC        = None      # path to a folds.csv from an earlier run: every model then validates on the same studies\nLABEL_BLEND     = True      # also test \"average of lexicon and label table\" as a target candidate (adopted only if it beats the lexicon by > 0.01 on gold)\nPHASE_A_SELECT  = \"val\"     # \"val\" = original (best epoch chosen on the validation fold, slightly optimistic OOF); \"last\" = final epoch, honest OOF\nN_BOOT          = 300       # bootstrap resamples for the gold-AUC confidence interval\nprint(f\"PRESET={PRESET} | {TIMM_NAME} | img {IMG_SIZE} | phase B {RUN_PHASE_B} | seed {SEED}\")\n\n# Record the settings the checkpoints depend on. The prediction notebook reads this file so that\n# preprocessing, image size and head settings cannot silently differ from training.\nimport os, json\n_CFG_KEYS = [\"PRESET\", \"MODEL_DIR\", \"TIMM_NAME\", \"IMG_SIZE\", \"CROP_MM\", \"SLOT_DEFS\", \"K_PER_SLOT\", \"FP_FIELDS\",\n             \"HEAD_HIDDEN\", \"PRIOR_W\", \"MAX_MIX\", \"LAB_MIX\", \"TTA\", \"SEED\"]\nos.makedirs(WORK, exist_ok=True)\njson.dump({k: globals()[k] for k in _CFG_KEYS}, open(f\"{WORK}/train_config.json\", \"w\"), indent=1)\n","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:07:11.727753Z","iopub.status.busy":"2026-09-19T06:07:11.726792Z","iopub.status.idle":"2026-09-19T06:07:11.738453Z","shell.execute_reply":"2026-09-19T06:07:11.737546Z"},"papermill":{"duration":0.018365,"end_time":"2026-09-19T06:07:11.740171+00:00","exception":false,"start_time":"2026-09-19T06:07:11.721806+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"49f55a14","cell_type":"markdown","source":"### 2 — Environment","metadata":{"papermill":{"duration":0.004529,"end_time":"2026-09-19T06:07:11.749558+00:00","exception":false,"start_time":"2026-09-19T06:07:11.745029+00:00","status":"completed"},"tags":[]}},{"id":"98014ad6","cell_type":"code","source":"import os, re, glob, time, unicodedata, warnings\nfrom concurrent.futures import ThreadPoolExecutor\nimport numpy as np, pandas as pd, cv2, pydicom\nimport torch, torch.nn as nn, torch.nn.functional as F, timm\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm.layers import resample_abs_pos_embed\nfrom timm.utils import ModelEmaV3\nfrom sklearn.feature_extraction.text import HashingVectorizer\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.linear_model import LogisticRegression\nwarnings.filterwarnings(\"ignore\")\n\nnp.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\ncv2.setNumThreads(0)\n\nDEV, GPU_IDS = \"cuda\", [0, 1]\nTRAIN_IMG, TEST_IMG = f\"{DATA_DIR}/train_series\", f\"{DATA_DIR}/test_series\"\nN_TOK    = sum(K_PER_SLOT)\nTOK2SLOT = np.concatenate([np.full(k, i, dtype=np.int64) for i, k in enumerate(K_PER_SLOT)])\nT0 = time.time()\nprint(f\"{N_TOK} tok @ {IMG_SIZE}px = {CROP_MM/IMG_SIZE:.3f} mm/px | {BATCH_STUDIES*N_TOK} img/step \"\n      f\"| STAGE={STAGE}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:07:11.759331Z","iopub.status.busy":"2026-09-19T06:07:11.758563Z","iopub.status.idle":"2026-09-19T06:07:30.378112Z","shell.execute_reply":"2026-09-19T06:07:30.377036Z"},"papermill":{"duration":18.626072,"end_time":"2026-09-19T06:07:30.379979+00:00","exception":false,"start_time":"2026-09-19T06:07:11.753907+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"3d1be254","cell_type":"markdown","source":"### 3 — Tables","metadata":{"papermill":{"duration":0.004182,"end_time":"2026-09-19T06:07:30.388653+00:00","exception":false,"start_time":"2026-09-19T06:07:30.384471+00:00","status":"completed"},"tags":[]}},{"id":"8d2fb442","cell_type":"code","source":"sub_df   = pd.read_csv(f\"{DATA_DIR}/sample_submission.csv\")\ntrain_df = pd.read_csv(f\"{DATA_DIR}/train.csv\")\ntest_df  = pd.read_csv(f\"{DATA_DIR}/test.csv\")\ntr_ser   = pd.read_csv(f\"{DATA_DIR}/train_series.csv\")\nte_ser   = pd.read_csv(f\"{DATA_DIR}/test_series.csv\")\n\nID, SER, REPORT = \"StudyInstanceUID\", \"SeriesInstanceUID\", \"Report\"\nLABELS = [c for c in sub_df.columns if c != ID]\nNL = len(LABELS)\ntrain_df[\"is_gold\"] = train_df[LABELS].notna().all(axis=1).astype(int)\n\nSERPATH, DCM_GLOB = \"{img}/{sid}/{ser}\", \"*.dcm\"\n\n_r0 = tr_ser.iloc[0]\n_d0 = SERPATH.format(img=TRAIN_IMG, sid=_r0[ID], ser=_r0[SER])\n_f0 = sorted(os.listdir(_d0))[:2]\nprint(\"filenames:\", [f[:34] + \"...\" for f in _f0])\nprint(\"  -> SOP UIDs carry no spatial order; slices are sorted by ImagePositionPatient in cell 9\")\nprint(f\"{len(train_df)} studies | {int(train_df.is_gold.sum())} gold | {NL} labels\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:07:30.39851Z","iopub.status.busy":"2026-09-19T06:07:30.398013Z","iopub.status.idle":"2026-09-19T06:07:30.681023Z","shell.execute_reply":"2026-09-19T06:07:30.679909Z"},"papermill":{"duration":0.289775,"end_time":"2026-09-19T06:07:30.682723+00:00","exception":false,"start_time":"2026-09-19T06:07:30.392948+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8c25f05a","cell_type":"markdown","source":"### 4 — Series metadata","metadata":{"papermill":{"duration":0.004077,"end_time":"2026-09-19T06:07:30.690872+00:00","exception":false,"start_time":"2026-09-19T06:07:30.686795+00:00","status":"completed"},"tags":[]}},{"id":"79f82d10","cell_type":"code","source":"ser_all = pd.concat(([tr_ser.assign(split=\"train\")] if STAGE == \"train\" else [])\n                    + [te_ser.assign(split=\"test\")], ignore_index=True)\nMETA = f\"{WORK}/meta_{STAGE}.parquet\"\n\ndef read_hdr(args):\n    sid, ser, sp = args\n    d0 = SERPATH.format(img=(TRAIN_IMG if sp == \"train\" else TEST_IMG), sid=sid, ser=ser)\n    files = sorted(glob.glob(os.path.join(d0, DCM_GLOB)))\n    if not files:\n        return None\n    try:\n        d = pydicom.dcmread(files[len(files) // 2], stop_before_pixels=True, force=True)\n    except Exception:\n        return None\n    iop, plane = getattr(d, \"ImageOrientationPatient\", None), \"unknown\"\n    if iop is not None and len(iop) == 6:\n        v = np.abs(np.cross([float(x) for x in iop[:3]], [float(x) for x in iop[3:]]))\n        plane = [\"sagittal\", \"coronal\", \"axial\"][int(np.argmax(v))]\n    return dict(\n        split=sp, sid=sid, ser=ser, n=len(files), plane=plane, dir=d0,\n        desc=str(getattr(d, \"SeriesDescription\", \"\")),\n        lat=str(getattr(d, \"Laterality\", \"\") or getattr(d, \"ImageLaterality\", \"\")),\n        fp=\"|\".join(str(getattr(d, a, \"\")) for a in FP_FIELDS),\n        TR=float(getattr(d, \"RepetitionTime\", 0) or 0),\n        TE=float(getattr(d, \"EchoTime\", 0) or 0),\n        TI=float(getattr(d, \"InversionTime\", 0) or 0))\n\nif not os.path.exists(META):\n    args = list(zip(ser_all[ID], ser_all[SER], ser_all.split))\n    with ThreadPoolExecutor(max_workers=16) as ex:\n        rows = [r for r in ex.map(read_hdr, args, chunksize=64) if r is not None]\n    pd.DataFrame(rows).to_parquet(META, index=False)\n    print(f\"headers: {len(rows)}/{len(args)} series | {(time.time()-T0)/60:.1f}m\")\nmeta = pd.read_parquet(META)\n\nFS_KW = r\"(stir|fs\\b|fatsat|fat[_\\- ]?sat|spair|spir|_fs|t2|tirm)\"\nmeta[\"fluid\"] = meta.desc.str.lower().str.contains(FS_KW, regex=True, na=False).astype(int)\nmeta.loc[(meta.TE > 40) & (meta.TR > 1500), \"fluid\"] = 1\nmeta.loc[(meta.TE < 20) & (meta.TR < 900),  \"fluid\"] = 0\nmeta.loc[meta.TI > 0, \"fluid\"] = 1\nfor col in ser_all.columns:\n    if \"fluid\" in col.lower():\n        meta[\"fluid\"] = meta.ser.map(ser_all.set_index(SER)[col].to_dict()).fillna(meta.fluid).astype(int)\n    if \"plane\" in col.lower():\n        meta[\"plane\"] = meta.ser.map(ser_all.set_index(SER)[col].astype(str).str.lower().to_dict()).fillna(meta.plane)\n\nmeta[\"slot\"] = -1\nfor i, (pl, fl) in enumerate(SLOT_DEFS):\n    meta.loc[(meta.plane == pl) & (meta.fluid == fl), \"slot\"] = i\n\nmeta[\"L\"] = meta.lat.str.upper().str[:1]\nmeta.loc[~meta.L.isin([\"L\", \"R\"]), \"L\"] = \"\"\ndl = meta.desc.str.lower()\nmeta.loc[(meta.L == \"\") & dl.str.contains(r\"\\b(left|lt|links|izquierd|sol|gauche|linker)\\b\",\n         regex=True, na=False), \"L\"] = \"L\"\nmeta.loc[(meta.L == \"\") & dl.str.contains(r\"\\b(right|rt|rechts|derech|sag|droit|rechter)\\b\",\n         regex=True, na=False), \"L\"] = \"R\"\n\nsrows = []\nfor sid, g in meta[meta.split == \"train\"].groupby(\"sid\"):\n    lv = g.loc[g.L != \"\", \"L\"]\n    f = g.fp.mode().iloc[0]\n    if not re.search(r\"[A-Za-z0-9]\", f):\n        f = \"anon:\" + sid                    # no scanner tags -> cannot leak site, own group\n    srows.append(dict(sid=sid, laterality=lv.mode().iloc[0] if len(lv) else \"R\", fp=f))\nsmeta = pd.DataFrame(srows if srows else [{\"sid\": train_df[ID].iloc[0], \"laterality\": \"R\",\n                                           \"fp\": \"unk\"}]).rename(columns={\"sid\": ID})\n\nprint(f\"slots {[int((meta.slot == i).sum()) for i in range(4)]} | \"\n      f\"fingerprints {smeta.fp.nunique()}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:07:30.700515Z","iopub.status.busy":"2026-09-19T06:07:30.699959Z","iopub.status.idle":"2026-09-19T06:08:56.045798Z","shell.execute_reply":"2026-09-19T06:08:56.044829Z"},"papermill":{"duration":85.352535,"end_time":"2026-09-19T06:08:56.04752+00:00","exception":false,"start_time":"2026-09-19T06:07:30.694985+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"711109ab","cell_type":"markdown","source":"### 5-6 — Weak labels\n\nOnly 58 of 4,407 studies are labelled, so targets are manufactured from radiology reports in nine\nlanguages. **This extractor is taken from the public notebook \"Bend the Knee to the Dinosaurs\"**\n(credits there: Pilkwang, Sofia Anjenje, Antoine G., prvsiyan, Marwan Mahmoud, dreaddevelopment /\nRoman Tamrazov, renta0426, Anvith Pothula).\n\nMeasured on the 58 gold studies against the hand-written lexicon it replaces:\n\n| | ours | theirs | |\n|---|---|---|---|\n| Medial OA | 0.689 | 0.891 | +0.202 |\n| Lateral OA | 0.626 | 0.813 | +0.188 |\n| Contusion | 0.696 | 0.855 | +0.159 |\n| PF OA | 0.707 | 0.808 | +0.100 |\n| Baker's | 0.837 | 0.924 | +0.087 |\n| **macro** | **0.774** | **0.862** | **+0.088** |\n\nIt wins where the hand-written version was weakest: compartment-level cartilage findings, which\nneed a side term and a cartilage term related across a window rather than matched as a phrase.\nIt also emits a per-finding `__conf`, which cell 7 uses as a second calibration feature.\n\nGiven the measured relationship LB ≈ 0.92 × target + 0.042, this moves the projection from\n~0.754 to ~0.835.\n","metadata":{"papermill":{"duration":0.004344,"end_time":"2026-09-19T06:08:56.056238+00:00","exception":false,"start_time":"2026-09-19T06:08:56.051894+00:00","status":"completed"},"tags":[]}},{"id":"7ecc9f7c","cell_type":"code","source":"import re\nimport unicodedata\nTARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\n_PRE = str.maketrans({'ı': 'i', 'İ': 'i', 'I': 'i', 'ß': 'ss', 'đ': 'd', 'Đ': 'd', 'ø': 'o', 'Ø': 'o', 'æ': 'ae', 'Æ': 'ae'})\n\ndef normalize(text: str) -> str:\n    if not isinstance(text, str):\n        return ''\n    text = text.translate(_PRE).lower()\n    text = unicodedata.normalize('NFKD', text)\n    text = ''.join((ch for ch in text if not unicodedata.combining(ch)))\n    text = text.replace('\\xad', '')\n    text = re.sub('[_\\\\-/\\\\\\\\]+', ' ', text)\n    text = re.sub('[ \\\\t]+', ' ', text)\n    return text\n_SENT_SPLIT = re.compile('(?<=[.;!?])\\\\s+|\\\\n+')\n\ndef unwrap(text: str) -> str:\n    if not isinstance(text, str):\n        return ''\n    out = []\n    for line in text.split('\\n'):\n        s = line.strip()\n        if out and out[-1] and (not re.search('[.;:!?>*•]$', out[-1])) and (len(out[-1].split()) >= 4) and s and (not s[:1].isupper()):\n            out[-1] = out[-1] + ' ' + s\n        else:\n            out.append(s)\n    return '\\n'.join(out)\n\ndef clauses(text: str):\n    norm = normalize(unwrap(text) if FEATURES['unwrap'] else text)\n    raw = [c.strip() for c in _SENT_SPLIT.split(norm) if c and c.strip()]\n    merged = []\n    for i, c in enumerate(raw):\n        if c.endswith(':') and len(c.split()) <= 14 and (i + 1 < len(raw)):\n            merged.append(c + ' ' + raw[i + 1])\n        merged.append(c)\n    out = []\n    for c in merged:\n        out.append(c)\n        if len(c.split()) > 25:\n            out.extend((p.strip() for p in c.split(',') if len(p.split()) > 2))\n    return out\nFEATURES = {'unwrap': True, 'directional_negation': True, 'oa_inherit': True, 'graded_pathology': True, 'synovitis_backoff': True}\n\ndef _rx(*alts: str) -> re.Pattern:\n    return re.compile('|'.join(alts))\nPRE_NEG = _rx('\\\\bno\\\\b', '\\\\bnot\\\\b', '\\\\bwithout\\\\b', '\\\\bnegative for\\\\b', '\\\\babsence\\\\b', '\\\\bno evidence\\\\b', '\\\\bfree of\\\\b', '\\\\bnone\\\\b', '\\\\bneither\\\\b', '\\\\bnor\\\\b', '\\\\bsin\\\\b', '\\\\bno hay\\\\b', '\\\\bausencia\\\\b', '\\\\bausentes?\\\\b', '\\\\bno se\\\\b', '\\\\bpas de\\\\b', '\\\\bsans\\\\b', '\\\\baucune?\\\\b', '\\\\bgeen\\\\b', '\\\\bzonder\\\\b', '\\\\bniet\\\\b', '\\\\bkeine?[nmrs]?\\\\b', '\\\\bohne\\\\b', '\\\\bnicht\\\\b', '\\\\bkein\\\\b', '\\\\bnema\\\\b', '\\\\bbez\\\\b', '\\\\bnisu\\\\b', '\\\\bnije\\\\b', '\\\\bδεν\\\\b', '\\\\bχωρις\\\\b', 'ουδεν', '\\\\bουτε\\\\b', '\\\\bбез\\\\b', '\\\\bне\\\\b', 'липсва', '\\\\bняма\\\\b')\nPOST_NEG = _rx('\\\\byok\\\\b', '\\\\byoktur\\\\b', 'izlenmemekte', 'saptanmadi', '\\\\bdegil\\\\b', 'gozlenmemekte', 'mevcut degil', 'eslik etmiyor', '\\\\bizlenmedi\\\\b', 'izlenmemistir', 'saptanmamistir', 'gorulmemistir', '\\\\bnema znakova\\\\b', 'bez znakova')\nNEGATION = _rx(PRE_NEG.pattern, POST_NEG.pattern, '\\\\bunremarkable\\\\b')\nNEG_WINDOW = 90\n\ndef _negated(clause: str, start: int, end: int) -> bool:\n    for m in PRE_NEG.finditer(clause):\n        if m.end() <= start and start - m.end() <= NEG_WINDOW:\n            if not re.search('\\\\b(but|however|ancak|fakat|pero|maar|aber|no i|ali|ωστοσο|αλλα|но)\\\\b', clause[m.end():start]):\n                return True\n    for m in POST_NEG.finditer(clause):\n        if m.start() >= end and m.start() - end <= NEG_WINDOW:\n            return True\n    return False\nNORMALITY = _rx('\\\\bnormal', '\\\\bintact\\\\b', '\\\\bpreserved\\\\b', '\\\\bwithin normal limits\\\\b', 'limites normales', '\\\\bconservad', '\\\\bintegr', '\\\\bnormales\\\\b', '\\\\bdoga(l|ll)\\\\b', 'korunmus', '\\\\bnormaldir\\\\b', 'olagan', '\\\\buredn', '\\\\bocuvan', '\\\\bodrzan', '\\\\bintakt', '\\\\bprimjeren', '\\\\bodrzanog kontinuiteta', '\\\\bodržan', 'φυσιολογικ', 'ακεραι', 'δεν παρατηρουνται', 'δεν σημειωνονται', 'unauffallig', 'regelrecht', '\\\\bo\\\\.?b\\\\.?\\\\b', 'нормал', 'запазен', 'съхранен', '\\\\bбез особености\\\\b', 'интактн', '\\\\bgaaf\\\\b', '\\\\bnormaal\\\\b')\nNORMAL_PHRASE = _rx('\\\\bsin alteracion', '\\\\bsin cambios\\\\b', '\\\\bsin particularidad', '\\\\bsin hallazgos\\\\b', '\\\\bsin lesion', '\\\\bsin signos de (rotura|lesion)', '\\\\bcontinu[oa]s?\\\\b', '\\\\bcontinuidad conservada\\\\b', '\\\\bno abnormalit', '\\\\bno significant abnormalit', '\\\\bunremarkable\\\\b', '\\\\bno evidence of (tear|injury|abnormalit)', '\\\\bohne auffalligkeit', '\\\\bkein nachweis\\\\b', '\\\\bohne befund\\\\b', '\\\\bgeen afwijking', '\\\\bzonder afwijking', '\\\\bsans anomalie', \"\\\\bpas d[e']anomalie\", '\\\\bbez osobitosti\\\\b', '\\\\bbez znakova (rupture|lezije)\\\\b', '\\\\bbez patoloskih\\\\b', 'χωρις αλλοιωσ', 'χωρις παθολογ', 'δεν παρατηρουνται (αξιολογα|παθολογ)', '\\\\bбез особености\\\\b', '\\\\bбез патологич', '\\\\bбез данни за\\\\b', '\\\\bozel bir ozellik yok', '\\\\bpatolojik bulgu (yok|izlenmemis)')\nUNCERTAIN = _rx('\\\\bpossible\\\\b', '\\\\bprobable\\\\b', '\\\\bsuspicious\\\\b', '\\\\bsuspected?\\\\b', 'cannot (be )?exclude', '\\\\bmay\\\\b', '\\\\bquestionable\\\\b', '\\\\bequivocal\\\\b', '\\\\br/o\\\\b', '\\\\bdd\\\\b', '\\\\blikely\\\\b', '\\\\bsuggest', '\\\\bcompatible with\\\\b', '\\\\bposible\\\\b', 'sin criterios categoricos', '\\\\bdudos', '\\\\bsugier', '\\\\bmuhtemel\\\\b', '\\\\bolasi\\\\b', '\\\\bsupheli\\\\b', '\\\\bizlenim', '\\\\bdusundur', '\\\\bmoguce\\\\b', '\\\\bvjerojatno\\\\b', '\\\\bsumnja\\\\b', '\\\\bmoze odgovarati\\\\b', 'πιθαν', 'υποπτ', '\\\\bmoglich', '\\\\bverdachtig', '\\\\bfraglich', '\\\\bv\\\\.?a\\\\.?\\\\b', '\\\\bwohl\\\\b', '\\\\bвъзможно\\\\b', '\\\\bвероятно\\\\b', 'суспект', '\\\\bmogelijk\\\\b', '\\\\bverdacht\\\\b')\n\nTEAR = _rx('\\\\btear', '\\\\btorn\\\\b', '\\\\brupture', '\\\\bdisruption\\\\b', 'discontinuit', '\\\\bavuls', '\\\\bmacerat', '\\\\bbuckethandle\\\\b', 'bucket handle', '\\\\brotura\\\\b', '\\\\broturas\\\\b', '\\\\bruptura', '\\\\bdesgarro', '\\\\broto\\\\b', '\\\\bdechirure', '\\\\bdechire', '\\\\bscheur', '\\\\bruptuur', 'gescheurd', '\\\\briss\\\\b', 'einriss', '\\\\bruptur', 'zerreiss', '\\\\blasion', '\\\\bausriss', '\\\\byirtik', '\\\\byirtig', '\\\\bkopma\\\\b', 'butunluk kaybi', '\\\\brupturu\\\\b', 'devamsizlik', '\\\\brupture\\\\b', '\\\\bdevamliligi secilememis', '\\\\bpuknuce', '\\\\bprekid\\\\b', '\\\\bpukotin', '\\\\bruptur', 'ρηξη', 'ρηξις', 'ρηγμα', 'ασυνεχεια', 'руптура', 'разкъсв', 'разрив', 'скъсв', '\\\\bлезия\\\\b')\nDEGEN = _rx('degenerat', '\\\\bmucoid\\\\b', '\\\\bmyxoid\\\\b', '\\\\bfray', '\\\\bfissur', 'dejeneratif', '\\\\bmukoid\\\\b', 'degenerativn', 'εκφυλ', 'дегенерат', '\\\\bμυξοειδ', '\\\\bμυξωδ', '\\\\bmeniskopat', '\\\\bmeniscopath', '\\\\bmuco ?ide\\\\b', 'aufgefasert', '\\\\bdejenerasyon\\\\b')\nINJURY = _rx('\\\\binjur', '\\\\bsprain', '\\\\blesion', '\\\\blasion', '\\\\bedema\\\\b', '\\\\boedema\\\\b', '\\\\bodem\\\\b', '\\\\bedem\\\\b', '\\\\bοιδημα', '\\\\bодем', '\\\\bедем', '\\\\bstrain\\\\b', '\\\\bhigh signal\\\\b', '\\\\bsignal alteration\\\\b', '\\\\bhiperintens', '\\\\bhyperintens', 'aumento de senal', 'alteracion de senal', 'cambio de senal', '\\\\bsignalanhebung', '\\\\bsignalalteration', 'verhoogd signaal', 'sinyal artis', 'αυξημενο σημα', 'повишен сигнал', '\\\\besguince\\\\b', '\\\\bthicken', '\\\\bzadebljanje\\\\b', '\\\\bverdikking\\\\b', '\\\\bdistenzij', '\\\\blaksite\\\\b', '\\\\blaxity\\\\b', '\\\\bpartial\\\\b', '\\\\bparcijaln', '\\\\bparcial', '\\\\bpartiel', '\\\\bpartiell')\n_GRADE_RX = re.compile('(?:grade|grad|grado|grau|derece|stupnja|stupanj|βαθμ|степен|icrs|outerbridge)[\\\\s:]*(?:grade\\\\s*)?([1-4]|iv|iii|ii|i)\\\\b')\n_ROMAN = {'i': 1, 'ii': 2, 'iii': 3, 'iv': 4}\n\ndef _grade_of(clause: str):\n    best = None\n    for m in _GRADE_RX.finditer(clause):\n        v = m.group(1)\n        n = _ROMAN.get(v, None) if not v.isdigit() else int(v)\n        if n is not None and (best is None or n > best):\n            best = n\n    return best\nANAT = {'ACL': _rx('anterior cruciate', '\\\\bacl\\\\b', 'cruzado anterior', '\\\\blca\\\\b', 'croise anterieur', 'voorste kruisband', '\\\\bvkb\\\\b', 'vorderes kreuzband', 'vorderen kreuzband', 'vordere kreuzband', 'on capraz', '\\\\bocb\\\\b', 'anterior capraz', 'prednji krizni', 'prednjeg krizn', 'προσθι[οα][^ ]* χιαστ', 'προσθιου χιαστου', 'χιαστο[^ ]* συνδεσμ', '\\\\bχιαστ\\\\w*', 'предна кръстна', 'предната кръстна', 'предна кръста', 'cruciate ligaments', 'ligamentos cruzados', 'ligaments croises', 'kruisbanden', 'kreuzbander', 'capraz baglar', 'krizn[a-z]* ligament[a-z]*', 'χιαστοι συνδεσμ', 'χιαστων συνδεσμ', 'кръстните връзки', 'кръстни връзки'), 'MCL': _rx('medial collateral', '\\\\bmcl\\\\b', 'tibial collateral', 'colateral medial', 'colateral interno', '\\\\blcm\\\\b', 'collateral medial', 'collateral interne', 'mediale collaterale', 'binnenband', '\\\\b(mediale|laterale) banden\\\\b', '\\\\bcollaterale banden\\\\b', 'innenband', 'mediales? kollateral', '\\\\bic yan bag', 'medial kollateral', '\\\\biyb\\\\b', 'medyal kollateral', 'medijalni kolateraln', 'medijalnog kolateraln', 'εσω πλαγι', 'εσωτερικο πλαγι', '\\\\bπλαγι\\\\w* συνδεσμ', '\\\\bπλαγιοι\\\\b', 'медиален колатерал', 'вътрешна странична', '\\\\bколатерал\\\\w*', '\\\\bcolaterales\\\\b', '\\\\bcollateraux\\\\b', '\\\\bcollateralen\\\\b', '\\\\bkolateralni\\\\b', 'collateral ligaments', 'ligamentos colaterales', 'ligaments collateraux', 'collaterale banden', 'kollateralbander', 'seitenbander', 'yan baglar', 'kolateraln[a-z]* ligament[a-z]*', 'πλαγιοι συνδεσμ', 'πλαγιων συνδεσμ', 'колатерални връзки', 'страничните връзки'), 'Medial Meniscus': _rx('medial meniscus', '\\\\bmm\\\\b(?= tear)', 'medial menisc', 'menisco medial', 'menisco interno', 'menisque medial', 'menisque interne', 'mediale meniscus', 'binnenmeniscus', 'innenmeniskus', 'medialen? meniskus', 'innenmeniskushinterhorn', 'medyal menisk', '\\\\bic menisk', 'medijalni meniskus', 'medijalnog meniskusa', 'medijalnom meniskusu', 'medijaln\\\\w* menisk\\\\w*', '\\\\bmedijalnog meniska\\\\b', 'medijalni menisk', 'εσω μηνισκ', 'μηνισκ[^ ]* του εσω', 'εσω διαμερισμα[^.]{0,40}μηνισκ', 'медиалния менискус', 'медиален менискус', 'вътрешния менискус', 'oba meniska', 'both menisci', 'ambos meniscos', 'beide menisci', 'her iki menisku', 'amfoteroi\\\\w* mhnisk', 'αμφοτερ\\\\w* μηνισκ', 'двата менискуса', 'medial (and|&) lateral menisc'), 'Lateral Meniscus': _rx('lateral meniscus', 'lateral menisc', 'menisco lateral', 'menisco externo', 'menisque lateral', 'menisque externe', 'laterale meniscus', 'buitenmeniscus', 'aussenmeniskus', 'lateralen? meniskus', 'aussenmeniskushinterhorn', 'lateral menisk', '\\\\bdis menisk', 'lateralni meniskus', 'lateralnog meniskusa', 'lateralnom meniskusu', 'lateraln\\\\w* menisk\\\\w*', '\\\\blateralnog meniska\\\\b', 'εξω μηνισκ', 'μηνισκ[^ ]* του εξω', 'εξω διαμερισμα[^.]{0,40}μηνισκ', 'латералния менискус', 'латерален менискус', 'външния менискус', 'oba meniska', 'both menisci', 'ambos meniscos', 'beide menisci', 'her iki menisku', 'αμφοτερ\\\\w* μηνισκ', 'двата менискуса', 'medial (and|&) lateral menisc')}\nOA_EVIDENCE = _rx('osteoarthrit', '\\\\barthros', '\\\\bgonarthros', '\\\\bosteoarthros', 'chondropath', 'chondromalac', 'condropat', 'condromalac', '\\\\bchondros', '\\\\bchondrosis\\\\b', 'chondral (loss|defect|ulcer|thinning|injury|fissur|wear)', 'cartilage (loss|thinning|defect|fissur|wear|damage|heterogeneity|irregularit)', '(loss|thinning|fissur|defect|ulcer|erosion|denudation) of[^.]{0,20}cartilage', 'articular cartilage[^.]{0,30}(loss|thin|fissur|defect|erosion|wear|irregular)', 'osteophyt', 'osteofit', 'osteofyt', 'osteofito', 'osteophyten', 'spurring', 'joint space narrowing', 'pinzamiento articular', 'reduced joint space', 'kikirdak kayb', 'kikirdak incelme', 'kondropati', 'kondral', 'kikirdak dejener', 'eklem aralig\\\\w* daral', 'eklem mesafesi daral', 'kikirdak kalinlig\\\\w* azal', 'kraakbeen', 'gonartrose', 'artrose', '\\\\bknorpel', 'arthrose', 'gonarthrose', 'hrskavic', 'hondromalac', 'artroz', 'osteoartrit', 'artrotsk', 'artrotick', '\\\\boa promjen', '\\\\boa\\\\b', 'degenerativne promjene hrskav', 'χονδρ[^ ]*παθ', 'αρθριτ', 'αρθρωσ', 'οστεοφυτ', 'χονδρομαλακ', 'αρθρικου χονδρου', 'εξαλειψη του αρθρικου χονδρου', 'διαβρωση του αρθρικου χονδρ', 'λεπτυνση[^.]{0,30}χονδρ', 'φθορα[^.]{0,20}χονδρ', 'артроз', 'хондропат', 'остеофит', 'хрущял[^.]{0,40}(изтън|увред|дефект|липс)', 'изтъняване[^.]{0,30}хрущял', 'хондромалац', 'ulcera[s]? condral', 'cartilago[^.]{0,25}(perdida|adelgaz)', 'icrs grade', 'icrs\\\\b', 'outerbridge', '\\\\bdenudation\\\\b', 'denudacij', 'erozivne promjene', '\\\\berosion of[^.]{0,20}cartilage', 'kraakbeenlijden', 'kraakbeenverlies')\nTF_SITE = _rx('compartment', 'compartimento', 'compartiment', 'kompartman', 'kompartiment', 'kompartment', 'odjelj', 'διαμερισμα', 'компартм', '\\\\bотдел', 'femorotibial', 'tibiofemoral', 'femoro tibial', 'femorotibiaal', 'femorotibijaln', 'феморотибиал', '\\\\bft zglob', 'tibiofemoraln', 'condyle', 'condilo', 'kondyl', 'kondil', 'condyl', 'κονδυλ', 'кондил', '\\\\bplateau', '\\\\bplato\\\\b', 'platillo', 'meseta', 'плато', 'tibiaplateau', 'tibijaln\\\\w* plato', 'tibyal plato', 'tibia plato', 'κνημιαι', 'μηριαι', 'weightbearing', 'weightbaring', 'zona de carga', 'dragende deel', 'agirlik tasiyan', '\\\\bfemur\\\\b', '\\\\btibia\\\\b', '\\\\bfemoral\\\\b', '\\\\btibial\\\\b', '\\\\bfemura\\\\b', '\\\\btibije\\\\b', '\\\\bmesarthrio\\\\b', 'μεσαρθριο')\nPF_SITE = _rx('patellofemoral', 'femoropatellar', 'femoropatelar', 'patelofemoral', 'retropatellar', 'retrorotulian', 'trochlea', 'troclea', 'troklea', 'trochlear', 'trohlej', 'τροχιλ', '\\\\bpatella', '\\\\bpatellar', 'rotulian', '\\\\brotula\\\\b', '\\\\bpatele\\\\b', 'patellofemoraal', 'femoropatellair', 'επιγονατιδ', 'μηροεπιγονατιδ', 'пател', 'феморопател', 'anterior compartment', 'compartimento anterior', 'prednj\\\\w* odjeljk', '\\\\bfp zglob', '\\\\bpf zglob', '\\\\bfaset', '\\\\bfacet', 'patellofemoraln')\nSIDE_MEDIAL = _rx('\\\\bmedial\\\\w*', '\\\\bmedyal\\\\w*', '\\\\bmedijaln\\\\w*', '\\\\bmediaal\\\\w*', '\\\\bmediale\\\\w*', '\\\\binterno\\\\b', '\\\\binterna\\\\b', '\\\\binternos\\\\b', '\\\\binterne\\\\b', '\\\\binnen\\\\w*', '\\\\bic\\\\b', '\\\\bunutarnj\\\\w*', '\\\\bεσω\\\\w*', '\\\\bεσωτερικ\\\\w*', '\\\\bмедиал\\\\w*', '\\\\bвътреш\\\\w*', '\\\\bbinnen\\\\w*', '\\\\bmediaal\\\\b', '\\\\bmediales?\\\\b')\nSIDE_LATERAL = _rx('\\\\blateral\\\\w*', '\\\\bexterno\\\\b', '\\\\bexterna\\\\b', '\\\\bexternos\\\\b', '\\\\bexterne\\\\b', '\\\\bdis\\\\b', '\\\\blateraln\\\\w*', '\\\\baussen\\\\w*', '\\\\bbuiten\\\\w*', '\\\\bεξω\\\\w*', '\\\\bεξωτερικ\\\\w*', '\\\\bлатерал\\\\w*', '\\\\bвъншн\\\\w*', '\\\\bvanjsk\\\\w*')\nSIDE_ANTERIOR = _rx('\\\\banterior\\\\w*', '\\\\bant\\\\b', '\\\\bon\\\\b', '\\\\bprednj\\\\w*', '\\\\bvorder\\\\w*', '\\\\bvoorste\\\\b', '\\\\bπροσθι\\\\w*', '\\\\bпредн\\\\w*', '\\\\banteriyor\\\\w*', '\\\\bavant\\\\b', '\\\\banterieur\\\\w*')\nGLOBAL_OA = _rx('tri ?compartment', 'all three compartment', 'global(ised)? (oa|osteoarthrit)', '\\\\bgonarthros', '\\\\bgonartros', '\\\\bgonarthrose', '\\\\bgonartrose', 'gonartro', 'goanrtrot', 'gonartrot', 'osteoarthritis of the knee', 'artrosis (de |)(la )?rodilla', 'knee osteoarthrit', '\\\\bdiz osteoartrit', '\\\\bgonartroz', 'artroza koljena', 'οστεοαρθριτιδα', 'αρθριτιδα του γονατος', 'εκφυλιστικη οστεοαρθριτ', 'артроза на колянната', 'гонартроз', 'degenerative joint disease', '\\\\bdjd\\\\b', 'three compartments', 'compartmens', 'compartments')\nDIRECT = {'Effusion': _rx('\\\\beffusion', 'joint fluid', 'intra ?articular fluid', '\\\\bhydrops\\\\b', '\\\\bhemarthros', '\\\\bhaemarthros', 'derrame articular', '\\\\bderrame\\\\b', 'liquido articular', 'hemartrosis', 'epanchement', 'gewrichtsvocht', '\\\\bvocht\\\\b', 'gewrichtseffusie', 'opzetting van suprapatell', 'gelenkerguss', '\\\\berguss\\\\b', 'gelenksergu', 'gelenksflussigkeit', 'eklem\\\\w* ic\\\\w* sivi', 'efuzyon', 'eklem sivisi', 'eklem mesafesinde sivi', 'sivi (miktari|artisi|birikimi)', 'sivi artis', '\\\\bsivi\\\\b[^.]{0,25}artmis', '\\\\bizljev', '\\\\bizliv', 'zglobn[^ ]* tekucin', '\\\\bhidrops\\\\b', 'αρθρικ[^ ]* υγρ', 'υγρου ενδαρθρικα', 'ενδαρθρικ[^ ]* υγρ', 'ποσοτητα υγρου', 'ενδαρθρικ', 'αρθρικη συλλογη', 'υγρο στην αρθρωση', 'υγρου στην αρθρωση', 'συλλογη υγρου', 'ενθαρθρικ', 'ставен излив', 'излив', 'ставна течност', 'синовиална течност'), 'Synovitis': _rx('synovit', 'sinovit', 'synovial (thickening|proliferation|hypertroph)', 'thicken\\\\w* synovial', 'hypertroph\\\\w* of the synovium', 'synoviale? (verdikking|proliferatie)', 'verdikkingen van (het )?synovium', 'synovialitis', 'synovialis(verdickung|proliferation)', 'reizsynovial', 'sinovijalitis', 'sinovitis', 'zadebljanje sinovij', 'proliferacij\\\\w* sinovij', 'sinovijaln\\\\w* proliferacij', 'υμενιτιδα', 'συνοβιτιδα', 'υμενικ[^ ]* υπερτροφ', 'αρθρικου υμεν', 'παχυνση[^.]{0,20}υμεν', 'υμενα', 'синовит', 'синовиал[^ ]* (задебел|пролифер)', '\\\\bpannus\\\\b', '\\\\bhoffit', 'sinovyal\\\\w* (kalinlas|proliferas)', 'sinovyal hipertrof', '\\\\bartrit\\\\b', '\\\\barthritis\\\\b'), \"Baker's\": _rx('baker', 'popliteal cyst', 'quiste popliteo', 'quistes popliteos', 'kyste poplite', 'popliteale? cyst', 'poplitealzyste', 'bakerzyste', 'popliteal kist', '\\\\bbakerova\\\\b', 'poplitealn[^ ]* cist', 'popliteal\\\\w* cist', 'κυστη baker', 'πολυχωρη συνοβιακη κυστη', 'κυστη του baker', 'συνοβιακη κυστη', 'κυστη τυπου baker', 'киста на бейкър', 'бейкърова киста', 'поплитеална киста', 'бекеров', 'gastrocnemio ?semimembranos', 'gastrocnemius semimembranosus burs'), 'Contusion': _rx('\\\\bcontusion', 'bone bruise', 'bone marrow (o?edema|contusion)', 'marrow o?edema', '\\\\bkontuz', 'medular bone o?edema', 'osseous contusion', 'contusion osea', 'edema oseo', 'edema de medula osea', 'contusiones oseas', 'oedeme osseux', 'contusion osseuse', 'botcontusie', 'botoedeem', 'beenmergoedeem', 'botmergoedeem', 'knochenmarkodem', 'knochenodem', 'knochenmarksodem', 'kontusion', 'kemik kontuzyonu', 'kemik iligi odemi', 'kemik odemi', 'kemik iliginde odem', 'kontuzyonel kemik', 'kemik iligi odemleri', 'kostani edem', 'edem kosti', 'kontuzij', 'kostane srzi[^.]{0,20}edem', 'οστεομυελικ[^ ]* οιδημα', 'οστικο οιδημα', 'μυελικο οιδημα', 'οστικο μωλωπ', 'костномозъчен едем', 'костен едем', 'контузионен', 'костно мозъчен едем'), 'Fracture': _rx('\\\\bfractur', '\\\\bfract\\\\b', '\\\\bfractura', '\\\\bfracturas\\\\b', '\\\\bfractuur', '\\\\bbreuk\\\\b', '\\\\bfraktur', '\\\\bbruch\\\\b', '\\\\bkirik\\\\b', '\\\\bkirigi\\\\b', '\\\\bkiri[kg]\\\\w*', '\\\\bprijelom', 'impresijsk[^ ]* fraktur', 'impaktcij', 'καταγμα', 'καταγματ', 'фрактур', 'счупван', 'фисур', 'insufficiency fracture', 'stress fracture', 'avulsion fracture', 'subchondral fracture', 'subkondral kiri', 'impaction (fracture|injury)', 'osteochondral (fracture|impaction)', '\\\\bsegond\\\\b', 'impactiefractuur', 'subchondrale impression', 'subchondraler? impress')}\nDECOY = {'Fracture': _rx('microfractur', '\\\\bfracture (risk|prophyla)'), \"Baker's\": _rx('meniscal cyst', 'quiste meniscal', 'parameniscal')}\nPAIRED = {'ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus'}\nOA_TARGETS = ['Medial OA', 'Lateral OA', 'PF OA']\nPLURAL_MENISCI = _rx('\\\\bmenisci\\\\b', '\\\\bmeniscos\\\\b', '\\\\bmenisques\\\\b', '\\\\bmenisken\\\\b', '\\\\bmeniskusi\\\\b', '\\\\bmenisk\\\\w*ler\\\\b', '\\\\bμηνισκοι\\\\b', '\\\\bμηνισκων\\\\b', '\\\\bменискуси\\\\b', '\\\\bменискусите\\\\b', '\\\\bmenisci\\\\w*\\\\b')\nANY_SIDE = _rx(SIDE_MEDIAL.pattern, SIDE_LATERAL.pattern)\nSTEM_MENISCUS = _rx('menisc\\\\w*', 'menisk\\\\w*', 'μηνισκ\\\\w*', 'мениск\\\\w*')\nSTEM_CRUCIATE = _rx('cruciate', 'cruzado', 'croise', 'kruisband', 'kreuzband', 'capraz bag\\\\w*', 'krizn\\\\w*', 'χιαστ\\\\w*', 'кръстн\\\\w*', '\\\\bacl\\\\b', '\\\\blca\\\\b', '\\\\bvkb\\\\b', '\\\\bocb\\\\b', '\\\\bacb\\\\b')\nSTEM_COLLATERAL = _rx('collateral\\\\w*', 'colateral\\\\w*', 'kollateral\\\\w*', 'collaterale\\\\w*', 'kolateraln\\\\w*', 'yan bag\\\\w*', 'πλαγι\\\\w*', 'колатерал\\\\w*', 'странич\\\\w*', 'innenband\\\\w*', 'binnenband\\\\w*', '\\\\bmcl\\\\b', '\\\\blcm\\\\b', '\\\\biyb\\\\b')\nSTEM_FRACTURE = _rx('fractur\\\\w*', 'fraktur\\\\w*', 'fractuur\\\\w*', '\\\\bfract\\\\b', 'kiri[kgğ]\\\\w*', 'prijelom\\\\w*', 'lom kosti', '\\\\bbreuk\\\\w*', '\\\\bbruch\\\\w*', 'καταγμα\\\\w*', 'καταγματ\\\\w*', 'фрактур\\\\w*', 'счупван\\\\w*', 'fisur\\\\w* (osea|oseas|kost)', 'fissur\\\\w* kost')\n\ndef _near(clause: str, stem_rx: re.Pattern, qual_rx: re.Pattern, window: int=55):\n    for m in stem_rx.finditer(clause):\n        lo = max(0, m.start() - window)\n        hi = min(len(clause), m.end() + window)\n        if qual_rx.search(clause[lo:hi]):\n            return True\n    return False\nSTEM_RULES = {'ACL': (STEM_CRUCIATE, SIDE_ANTERIOR), 'MCL': (STEM_COLLATERAL, SIDE_MEDIAL), 'Medial Meniscus': (STEM_MENISCUS, SIDE_MEDIAL), 'Lateral Meniscus': (STEM_MENISCUS, SIDE_LATERAL)}\n\nclass _Matcher:\n\n    def __init__(self, phrase_rx, stem=None, side=None, window=55):\n        self.phrase_rx = phrase_rx\n        self.stem = stem\n        self.side = side\n        self.window = window\n\n    def search(self, clause):\n        m = self.phrase_rx.search(clause)\n        if m is not None:\n            return m\n        if self.stem is not None and _near(clause, self.stem, self.side, self.window):\n            return self.stem.search(clause)\n        return None\nANAT_MATCH = {t: _Matcher(ANAT[t], *STEM_RULES[t]) for t in PAIRED}\nDIRECT_MATCH = {t: _Matcher(_rx(rx.pattern, STEM_FRACTURE.pattern) if t == 'Fracture' else rx) for t, rx in DIRECT.items()}\nSEV_LOW = _rx('\\\\bsmall\\\\b', '\\\\bminimal\\\\b', '\\\\btrace\\\\b', '\\\\bmild\\\\b', '\\\\bslight\\\\b', '\\\\btiny\\\\b', '\\\\bscant\\\\b', '\\\\bdiscrete\\\\b', '\\\\blow ?grade\\\\b', '\\\\bincipient\\\\b', '\\\\bleve\\\\b', '\\\\bminim', '\\\\bpeque', '\\\\bfina\\\\b', '\\\\bfino\\\\b', '\\\\bligero\\\\b', '\\\\bescaso\\\\b', '\\\\bdiscreto\\\\b', '\\\\bhafif\\\\b', '\\\\baz miktarda\\\\b', '\\\\bsilik\\\\b', '\\\\bmanj\\\\w*', '\\\\bblago\\\\b', '\\\\bdiskretn', '\\\\bmalo\\\\b', '\\\\bpocetn', '\\\\bgering', '\\\\bdiskret', '\\\\bkleine?r?\\\\b', '\\\\bwenig\\\\b', '\\\\bzarte?\\\\b', '\\\\bbeperkte?\\\\b', '\\\\bgeringe\\\\b', '\\\\bweinig\\\\b', '\\\\blichte?\\\\b', '\\\\blicht\\\\b', '\\\\bηπι', '\\\\bμικρ', '\\\\bελαχιστ', '\\\\bαρχομεν', '\\\\bминимал', '\\\\bлек', '\\\\bмалк', '\\\\bнеголям')\nSEV_HIGH = _rx('\\\\blarge\\\\b', '\\\\bmarked\\\\b', '\\\\bmassive\\\\b', '\\\\bsevere\\\\b', '\\\\bextensive\\\\b', '\\\\bmoderate\\\\b', '\\\\bgross\\\\b', '\\\\bsignificant\\\\b', '\\\\babundant\\\\b', '\\\\btense\\\\b', '\\\\bcomplete\\\\b', '\\\\bfull ?thickness\\\\b', '\\\\bhigh ?grade\\\\b', '\\\\badvanced\\\\b', '\\\\bmoderad', '\\\\bimportante\\\\b', '\\\\bsevera?\\\\b', '\\\\bmarcad', '\\\\bcuantios', '\\\\bespesor total\\\\b', '\\\\bcompleta?\\\\b', '\\\\bbelirgin\\\\b', '\\\\byaygin\\\\b', '\\\\bileri\\\\b', '\\\\bciddi\\\\b', '\\\\bbol\\\\b', '\\\\bkomplet', '\\\\bopsezan\\\\b', '\\\\bveliki\\\\b', '\\\\bizrazit', '\\\\bznacajn', '\\\\bumjeren', '\\\\buznapredoval', '\\\\bpotpun', '\\\\bkompleksn', '\\\\bausgepragt', '\\\\bdeutlich', '\\\\bmassiv', '\\\\bmassig', '\\\\bgross', '\\\\buitgebreid', '\\\\bgevorderd', '\\\\bveel\\\\b', '\\\\bmatige?\\\\b', '\\\\bvolledig', '\\\\bμετρι', '\\\\bμεγαλ', '\\\\bεκτεταμεν', '\\\\bευμεγεθ', '\\\\bσοβαρ', '\\\\bπληρη', '\\\\bголям', '\\\\bизразен', '\\\\bзначим', '\\\\bумерен', '\\\\bобилен', '\\\\bпълн')\nGRADE_HIGH = re.compile('grade?[ao]?\\\\s*(3|4|iii|iv)\\\\b|icrs grade (iii|iv|3|4)|stupnja iv|stupnja iii|\\\\bgrado (3|4)\\\\b|\\\\bgrad (3|4)\\\\b|\\\\bgrade (3|4)\\\\b')\nDEGENERATIVE_MARROW = _rx('subchondral', 'subcondral', 'subkondral', 'supkondraln', 'subchondraln', 'υποχονδρι', 'υπαρθρικ', 'субхондрал', 'subchondrale?', 'subartikuler', '\\\\bcyst', '\\\\bquist', '\\\\bzyste\\\\b', '\\\\bcistic', 'reactive', 'reactivo', 'degenerative', 'degenerativ', 'reaktiv', '\\\\bcisti\\\\b')\nTRAUMA = _rx('\\\\bbruise\\\\b', '\\\\bcontusion', '\\\\bkontuz', '\\\\btrauma', '\\\\bimpaction\\\\b', '\\\\bpivot shift\\\\b', '\\\\bkissing\\\\b', '\\\\bacute\\\\b', '\\\\bagudo\\\\b', '\\\\bakut', '\\\\bpivot kaymasi\\\\b', '\\\\bcontusion osseuse\\\\b', '\\\\bbone bruise\\\\b', '\\\\bbotcontusie\\\\b', '\\\\bконтузион', '\\\\bμωλωπ', '\\\\bkontuzij', '\\\\bimpaktcij', '\\\\bimpakcij', '\\\\bfall\\\\b', '\\\\binjury\\\\b', '\\\\bimpression\\\\b')\nSYNOVIAL_PROXY = _rx('bursit', 'burzit', '\\\\bbursa\\\\b[^.]{0,30}(fluid|distend|sivi|tekucin|opzetting)', 'suprapatellar (bursitis|effusion|recess)', 'suprapatellar bursa', 'suprapatellar bursada', 'suprapatelarno', 'suprapatellaire recessus', 'hoffa', 'hoffit', 'plica', 'plika', 'πλικα', 'fat pad[^.]{0,20}(edema|oedema)', 'kapsul', 'capsul', 'καψ', 'капсул', '\\\\bpannus\\\\b', '\\\\bsinov', '\\\\bsynov')\n\ndef _polarity(clause: str, span=None) -> str:\n    if UNCERTAIN.search(clause):\n        return 'uncertain'\n    if span is None or not FEATURES['directional_negation']:\n        if NEGATION.search(clause):\n            return 'negative'\n    elif _negated(clause, span[0], span[1]):\n        return 'negative'\n    if NORMALITY.search(clause):\n        if TEAR.search(clause) or GRADE_HIGH.search(clause):\n            return 'positive'\n        return 'negative'\n    return 'positive'\n\ndef _severity(clause: str) -> float:\n    high = SEV_HIGH.search(clause) is not None\n    low = SEV_LOW.search(clause) is not None\n    if high and (not low):\n        return 1.0\n    if low and (not high):\n        return 0.45\n    if high and low:\n        return 0.8\n    return 0.75\n\ndef _grade(n_pos, n_neg, n_unc, best):\n    if n_pos or n_unc:\n        score = min(0.97, 0.5 + 0.45 * best + 0.015 * min(n_pos, 3))\n        conf = min(1.0, 0.55 + 0.15 * n_pos)\n    elif n_neg:\n        score = max(0.04, 0.2 - 0.04 * n_neg)\n        conf = min(0.9, 0.45 + 0.12 * n_neg)\n    else:\n        score, conf = (0.28, 0.05)\n    return (score, conf)\n\ndef _paired_weight(clause: str, meniscus: bool) -> float:\n    g = _grade_of(clause) if FEATURES['graded_pathology'] else None\n    tear = TEAR.search(clause) is not None\n    if meniscus:\n        if tear:\n            base = 1.0\n        elif g is not None:\n            base = 0.95 if g >= 3 else 0.3\n        elif DEGEN.search(clause):\n            base = 0.35\n        else:\n            base = 0.45\n    elif tear:\n        base = 1.0\n    elif g is not None:\n        base = 0.85 if g >= 2 else 0.3\n    elif DEGEN.search(clause):\n        base = 0.4\n    else:\n        base = 0.55\n    if SEV_HIGH.search(clause) and (not SEV_LOW.search(clause)):\n        base = min(1.0, base * 1.2)\n    elif SEV_LOW.search(clause) and (not SEV_HIGH.search(clause)):\n        base *= 0.7\n    return base\n\ndef _score_paired(cls, tgt):\n    anat_rx = ANAT_MATCH[tgt]\n    path_rx = _rx(TEAR.pattern, DEGEN.pattern, INJURY.pattern)\n    meniscus = 'Meniscus' in tgt\n    n_pos = n_neg = n_unc = 0\n    best = 0.0\n    for c in cls:\n        hit = anat_rx.search(c)\n        if hit is None and meniscus and PLURAL_MENISCI.search(c) and (not ANY_SIDE.search(c)):\n            hit = PLURAL_MENISCI.search(c)\n        if hit is None:\n            continue\n        pm = path_rx.search(c)\n        if pm is None and _grade_of(c) is None:\n            if NORMAL_PHRASE.search(c) or (NORMALITY.search(c) and (not NEGATION.search(c))):\n                n_neg += 1\n            continue\n        span = (pm.start(), pm.end()) if pm is not None else None\n        pol = _polarity(c, span)\n        if pol == 'positive':\n            n_pos += 1\n            best = max(best, _paired_weight(c, meniscus))\n        elif pol == 'negative':\n            n_neg += 1\n        else:\n            n_unc += 1\n            best = max(best, 0.45 * _paired_weight(c, meniscus))\n    s, cf = _grade(n_pos, n_neg, n_unc, best)\n    return (s, cf, n_pos, n_neg)\n\ndef _score_clauses(cls, anat_rx, path_rx=None, decoy_rx=None, context_penalty=None, context_bonus=None):\n    n_pos = n_neg = n_unc = 0\n    best = 0.0\n    for c in cls:\n        m = anat_rx.search(c)\n        if not m:\n            continue\n        if decoy_rx is not None and decoy_rx.search(c):\n            continue\n        if path_rx is not None and (not path_rx.search(c)):\n            if NORMAL_PHRASE.search(c) or (NORMALITY.search(c) and (not NEGATION.search(c))):\n                n_neg += 1\n            continue\n        pol = _polarity(c, (m.start(), m.end()))\n        if pol == 'positive':\n            n_pos += 1\n            w = _severity(c)\n            if context_penalty is not None and context_penalty.search(c):\n                w *= 0.45\n            if context_bonus is not None and context_bonus.search(c):\n                w = min(1.0, w * 1.35)\n            best = max(best, w)\n        elif pol == 'negative':\n            n_neg += 1\n        else:\n            n_unc += 1\n            best = max(best, 0.3)\n    s, c = _grade(n_pos, n_neg, n_unc, best)\n    return (s, c, n_pos, n_neg)\n\ndef _score_oa(cls):\n    acc = {t: {'pos': 0, 'neg': 0, 'unc': 0, 'best': 0.0} for t in OA_TARGETS}\n    g_pos, g_neg, g_best = (0, 0, 0.0)\n    for c in cls:\n        m = OA_EVIDENCE.search(c)\n        if not m:\n            continue\n        pol = _polarity(c, (m.start(), m.end()))\n        sev = _severity(c)\n        tf_med = _near(c, TF_SITE, SIDE_MEDIAL, 45)\n        tf_lat = _near(c, TF_SITE, SIDE_LATERAL, 45)\n        pf = PF_SITE.search(c) is not None\n        hits = []\n        if tf_med:\n            hits.append('Medial OA')\n        if tf_lat:\n            hits.append('Lateral OA')\n        if pf:\n            hits.append('PF OA')\n        if not hits:\n            if pol == 'positive':\n                g_pos += 1\n                g_best = max(g_best, sev if GLOBAL_OA.search(c) else sev * 0.7)\n            elif pol == 'negative':\n                g_neg += 1\n            continue\n        for t in hits:\n            if pol == 'positive':\n                acc[t]['pos'] += 1\n                acc[t]['best'] = max(acc[t]['best'], sev)\n            elif pol == 'negative':\n                acc[t]['neg'] += 1\n            else:\n                acc[t]['unc'] += 1\n                acc[t]['best'] = max(acc[t]['best'], 0.3)\n    out = {}\n    for t in OA_TARGETS:\n        a = acc[t]\n        pos, neg, unc, best = (a['pos'], a['neg'], a['unc'], a['best'])\n        if not (pos or unc) and g_pos and FEATURES['oa_inherit']:\n            if neg:\n                score, conf = _grade(0, neg, 0, 0.0)\n                score = max(score, 0.35)\n                conf *= 0.7\n            else:\n                score, conf = _grade(g_pos, 0, 0, g_best * 0.92)\n                conf *= 0.75\n        else:\n            score, conf = _grade(pos, neg + g_neg, unc, best)\n        out[t] = (score, conf, pos, neg)\n    return out\n\ndef extract(report: str) -> dict:\n    cls = clauses(report)\n    out = {}\n    for tgt in PAIRED:\n        s, c, npos, nneg = _score_paired(cls, tgt)\n        out[tgt] = s\n        out[tgt + '__conf'] = c\n        out[tgt + '__npos'] = npos\n        out[tgt + '__nneg'] = nneg\n    for tgt, (s, c, npos, nneg) in _score_oa(cls).items():\n        out[tgt] = s\n        out[tgt + '__conf'] = c\n        out[tgt + '__npos'] = npos\n        out[tgt + '__nneg'] = nneg\n    for tgt in ('Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture'):\n        if tgt == 'Contusion':\n            s, c, npos, nneg = _score_clauses(cls, DIRECT_MATCH[tgt], None, DECOY.get(tgt), context_penalty=DEGENERATIVE_MARROW, context_bonus=TRAUMA)\n        else:\n            s, c, npos, nneg = _score_clauses(cls, DIRECT_MATCH[tgt], None, DECOY.get(tgt))\n        out[tgt] = s\n        out[tgt + '__conf'] = c\n        out[tgt + '__npos'] = npos\n        out[tgt + '__nneg'] = nneg\n    if FEATURES['synovitis_backoff'] and out['Synovitis__npos'] == 0 and (out['Synovitis__nneg'] == 0):\n        proxy = sum((1 for c in cls if SYNOVIAL_PROXY.search(c) and _polarity(c) == 'positive'))\n        eff = out['Effusion']\n        prior = 0.3 + 0.3 * max(0.0, (eff - 0.5) / 0.45) + 0.06 * min(proxy, 3)\n        out['Synovitis'] = min(0.72, prior)\n        out['Synovitis__conf'] = 0.18\n    return out\n\nsev  = np.zeros((len(train_df), NL), np.float32)\nconf = np.zeros((len(train_df), NL), np.float32)\nfor i, t in enumerate(train_df[REPORT].fillna(\"\").astype(str).values):\n    r = extract(t)\n    for j, L in enumerate(LABELS):\n        sev[i, j]  = float(r.get(L, 0))\n        conf[i, j] = float(r.get(L + \"__conf\", 0))\n    if i % 1500 == 0:\n        print(f\"  report {i}/{len(train_df)}\")\nprint(\"fired%:\", {L: round(float((sev[:, j] > 0).mean() * 100), 1) for j, L in enumerate(LABELS)})\n\n\n\n\n# Report labels. The public 0.941 notebooks read a table from data/derived/report_labels_v2.csv,\n# which is a private path — no such file is published, so there is nothing to auto-discover.\n# Point LABEL_TABLE at your own CSV (StudyInstanceUID + the 12 targets, optionally + \"__conf\")\n# if you ever build better labels; it is adopted only if it beats the lexicon on the gold set.\nLABEL_SOURCE = \"lexicon\"\n_gi = np.where(train_df[LABELS].notna().all(axis=1).values)[0]\n_yg = train_df.loc[train_df.is_gold == 1, LABELS].values.astype(float)\n_ok = [j for j in range(NL) if 0 < _yg[:, j].sum() < len(_yg)]\n_base = float(np.mean([roc_auc_score(_yg[:, j], sev[_gi, j]) for j in _ok]))\nprint(f\"lexicon target quality on {len(_gi)} gold: {_base:.4f}\")\n\nif LABEL_TABLE and not os.path.exists(LABEL_TABLE):\n    print(f\"!! LABEL_TABLE not found: {LABEL_TABLE}\")\n    LABEL_TABLE = None\n\nif LABEL_TABLE is None and LABEL_DIR:\n    if not os.path.isdir(LABEL_DIR):\n        print(f\"!! LABEL_DIR does not exist: {LABEL_DIR}\")\n        print(\"   /kaggle/input holds:\", sorted(os.listdir(\"/kaggle/input\")))\n    else:\n        _csv = sorted(glob.glob(f\"{LABEL_DIR}/**/*.csv\", recursive=True))\n        for _p in _csv:\n            _h = pd.read_csv(_p, nrows=1)\n            if ID in _h.columns and all(t in _h.columns for t in LABELS):\n                LABEL_TABLE = _p\n                break\n            if ID in _h.columns:\n                print(f\"   {os.path.basename(_p)} has the ID but is missing \"\n                      f\"{[t for t in LABELS if t not in _h.columns][:4]}\")\n        if LABEL_TABLE is None:\n            print(f\"!! no usable CSV in {LABEL_DIR}; files there: {[os.path.basename(x) for x in _csv][:8]}\")\n            print(f\"   a usable table needs '{ID}' plus all 12 target columns\")\n    print(\"label table:\", LABEL_TABLE)\n\n# Falling back to the lexicon silently would burn 8 hours training on weaker targets and say so\n# in one log line. Fail here instead.\nassert LABEL_TABLE or not (LABEL_DIR and REQUIRE_LABELS), (\n    f\"LABEL_DIR is set to {LABEL_DIR} but no usable label table was found. \"\n    \"Attach the dataset, fix the path, or set REQUIRE_LABELS=False to train on the lexicon.\")\n\nif LABEL_TABLE and os.path.exists(LABEL_TABLE):\n    _t = pd.read_csv(LABEL_TABLE).drop_duplicates(ID).set_index(ID)\n    _hit = train_df[ID].isin(_t.index).values\n    _s = sev.copy()\n    _s[_hit] = _t.loc[train_df.loc[_hit, ID], LABELS].values.astype(np.float32)\n    _a = float(np.mean([roc_auc_score(_yg[:, j], _s[_gi, j]) for j in _ok]))\n    print(f\"{os.path.basename(LABEL_TABLE)}: covers {_hit.mean()*100:.0f}%, gold {_a:.4f} ({_a-_base:+.4f})\")\n    _cand = {\"table\": _s}\n    if LABEL_BLEND:\n        _cand[\"average\"] = 0.5 * (sev + _s)\n    _score = {\"table\": _a}\n    if \"average\" in _cand:\n        _score[\"average\"] = float(np.mean([roc_auc_score(_yg[:, j], _cand[\"average\"][_gi, j]) for j in _ok]))\n        print(f\"  average of lexicon and table: gold {_score['average']:.4f} ({_score['average']-_base:+.4f})\")\n    _pick = max(_score, key=_score.get)\n    LABEL_SOURCE = \"lexicon\"\n    if _score[_pick] > _base + 0.01:\n        sev = _cand[_pick]\n        _cc = [t + \"__conf\" for t in LABELS]\n        if _pick == \"table\" and all(c in _t.columns for c in _cc):\n            conf[_hit] = _t.loc[train_df.loc[_hit, ID], _cc].values.astype(np.float32)\n        LABEL_SOURCE = _pick\n        print(\"  adopted\", _pick)\n    else:\n        print(\"  rejected, keeping the lexicon\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:08:56.066905Z","iopub.status.busy":"2026-09-19T06:08:56.066638Z","iopub.status.idle":"2026-09-19T06:09:20.216773Z","shell.execute_reply":"2026-09-19T06:09:20.215862Z"},"papermill":{"duration":24.157872,"end_time":"2026-09-19T06:09:20.218356+00:00","exception":false,"start_time":"2026-09-19T06:08:56.060484+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0d09b569","cell_type":"markdown","source":"### 7 — Soft targets and folds","metadata":{"papermill":{"duration":0.00453,"end_time":"2026-09-19T06:09:20.227469+00:00","exception":false,"start_time":"2026-09-19T06:09:20.222939+00:00","status":"completed"},"tags":[]}},{"id":"d9209f5c","cell_type":"code","source":"gold = np.where(train_df.is_gold.values == 1)[0]\nyg   = train_df.loc[train_df.is_gold == 1, LABELS].values.astype(np.float32)\n\nX_all = np.stack([sev, conf], -1)\nsoft = np.zeros_like(sev)\nfor j in range(NL):\n    soft[:, j] = np.clip(LogisticRegression(C=1.0, max_iter=1000)\n                         .fit(X_all[gold, j], yg[:, j])\n                         .predict_proba(X_all[:, j])[:, 1], 0.01, 0.97)\n\nok = [j for j in range(NL) if 0 < yg[:, j].sum() < len(yg)]\nprint(f\"targets vs gold  label {np.mean([roc_auc_score(yg[:,j], sev[gold,j]) for j in ok]):.4f}\"\n      f\"   conf {np.mean([roc_auc_score(yg[:,j], conf[gold,j]) for j in ok]):.4f}\"\n      f\"   (n=58; the calibrator sees these 58, so its own AUC is not a held-out number)\")\n\nsw = np.ones_like(soft)\nsoft[gold, :], sw[gold, :] = yg, GOLD_W\npc = np.clip(soft, 1e-6, 1 - 1e-6)\nFLOOR = float((-(pc * np.log(pc) + (1 - pc) * np.log(1 - pc)) * sw).mean())\n\ntxt_t = HashingVectorizer(n_features=64, analyzer=\"char_wb\", ngram_range=(3, 5),\n                          alternate_sign=False, norm=\"l2\") \\\n        .transform(train_df[REPORT].fillna(\"\").astype(str)).toarray().astype(np.float32)\n\ntrain_df = train_df.merge(smeta, on=ID, how=\"left\")\ntrain_df[\"fp\"] = train_df.fp.fillna(\"unk\")\ntrain_df[\"laterality\"] = train_df.laterality.fillna(\"R\")\nload, gload, assign = np.zeros(N_FOLDS), np.zeros(N_FOLDS), {}\ngsz, ggold = train_df.fp.value_counts(), train_df.groupby(\"fp\").is_gold.sum()\nfor g in gsz.index:\n    f = int(np.lexsort((gload, load))[0])\n    assign[g] = f; load[f] += gsz[g]; gload[f] += ggold.get(g, 0)\ntrain_df[\"fold\"] = train_df.fp.map(assign).astype(int)\nif FOLD_SRC and os.path.exists(FOLD_SRC):\n    _fs = pd.read_csv(FOLD_SRC).drop_duplicates(ID).set_index(ID)[\"fold\"]\n    assert train_df[ID].isin(_fs.index).all(), \"FOLD_SRC does not cover every training study\"\n    train_df[\"fold\"] = train_df[ID].map(_fs).astype(int)\n    print(f\"folds loaded from {FOLD_SRC}\")\nelif FOLD_SRC:\n    print(f\"!! FOLD_SRC not found ({FOLD_SRC}); using the notebook's own fold assignment\")\ntrain_df[[ID, \"fold\"]].to_csv(f\"{WORK}/folds.csv\", index=False)\nnp.save(f\"{WORK}/soft.npy\", soft)\nleak = sum(len(set(train_df.loc[train_df.fold == f, \"fp\"]) &\n               set(train_df.loc[train_df.fold != f, \"fp\"])) for f in range(N_FOLDS))\nprint(f\"loss floor {FLOOR:.3f} | fold sizes {train_df.groupby('fold').size().tolist()} \"\n      f\"| leaked {leak} | gold/fold {train_df.groupby('fold').is_gold.sum().tolist()}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:09:20.237541Z","iopub.status.busy":"2026-09-19T06:09:20.237305Z","iopub.status.idle":"2026-09-19T06:09:25.226621Z","shell.execute_reply":"2026-09-19T06:09:25.225853Z"},"papermill":{"duration":4.996879,"end_time":"2026-09-19T06:09:25.228666+00:00","exception":false,"start_time":"2026-09-19T06:09:20.231787+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a122df3e","cell_type":"markdown","source":"### 8 — Backbone","metadata":{"papermill":{"duration":0.00602,"end_time":"2026-09-19T06:09:25.241322+00:00","exception":false,"start_time":"2026-09-19T06:09:25.235302+00:00","status":"completed"},"tags":[]}},{"id":"29035bdc","cell_type":"code","source":"bb = timm.create_model(TIMM_NAME, pretrained=False, num_classes=0, img_size=IMG_SIZE, global_pool=\"\")\nEMB_DIM = 3 * bb.num_features      # CLS + mean-patch + top-k patch mean\ntgt, ntot = bb.state_dict(), len(bb.state_dict())\n\nwf = [f for f in sorted(os.listdir(MODEL_DIR)) if f.endswith((\".bin\", \".safetensors\", \".pth\", \".pt\"))]\nassert wf, f\"no weight file in {MODEL_DIR}: {sorted(os.listdir(MODEL_DIR))}\"\nCKPT = os.path.join(MODEL_DIR, max(wf, key=lambda f: os.path.getsize(os.path.join(MODEL_DIR, f))))\nif CKPT.endswith(\".safetensors\"):\n    from safetensors.torch import load_file\n    raw = load_file(CKPT)\nelse:\n    raw = torch.load(CKPT, map_location=\"cpu\", weights_only=False)\n    raw = raw.get(\"model\", raw.get(\"state_dict\", raw.get(\"teacher\", raw)))\nraw = {re.sub(r\"^(module\\.|backbone\\.|dinov2\\.)+\", \"\", k): v\n       for k, v in raw.items() if hasattr(v, \"shape\")}\n\nsd = {}\nif any(k.startswith(\"encoder.layer.\") for k in raw):\n    depth = 1 + max(int(k.split(\".\")[2]) for k in raw if k.startswith(\"encoder.layer.\"))\n    for i in range(depth):\n        p = f\"encoder.layer.{i}.\"\n        for part in (\"weight\", \"bias\"):\n            qkv = [raw.get(p + f\"attention.attention.{n}.{part}\") for n in (\"query\", \"key\", \"value\")]\n            if all(x is not None for x in qkv):\n                sd[f\"blocks.{i}.attn.qkv.{part}\"] = torch.cat(qkv, 0)\n        for a, b in [(\"attention.output.dense.weight\", \"attn.proj.weight\"),\n                     (\"attention.output.dense.bias\", \"attn.proj.bias\"),\n                     (\"layer_scale1.lambda1\", \"ls1.gamma\"), (\"layer_scale2.lambda1\", \"ls2.gamma\"),\n                     (\"norm1.weight\", \"norm1.weight\"), (\"norm1.bias\", \"norm1.bias\"),\n                     (\"norm2.weight\", \"norm2.weight\"), (\"norm2.bias\", \"norm2.bias\"),\n                     (\"mlp.fc1.weight\", \"mlp.fc1.weight\"), (\"mlp.fc1.bias\", \"mlp.fc1.bias\"),\n                     (\"mlp.fc2.weight\", \"mlp.fc2.weight\"), (\"mlp.fc2.bias\", \"mlp.fc2.bias\")]:\n            if p + a in raw:\n                sd[f\"blocks.{i}.{b}\"] = raw[p + a]\n    for a, b in [(\"embeddings.cls_token\", \"cls_token\"),\n                 (\"embeddings.position_embeddings\", \"pos_embed\"),\n                 (\"embeddings.patch_embeddings.projection.weight\", \"patch_embed.proj.weight\"),\n                 (\"embeddings.patch_embeddings.projection.bias\", \"patch_embed.proj.bias\"),\n                 (\"layernorm.weight\", \"norm.weight\"), (\"layernorm.bias\", \"norm.bias\")]:\n        if a in raw:\n            sd[b] = raw[a]\n    layout = f\"HF transformers, {depth} layers, q/k/v fused\"\nelse:\n    sd, layout = dict(raw), \"Meta original\"\n\nif \"pos_embed\" in sd and sd[\"pos_embed\"].shape[1] != bb.pos_embed.shape[1]:\n    gs = IMG_SIZE // bb.patch_embed.patch_size[0]\n    sd[\"pos_embed\"] = resample_abs_pos_embed(sd[\"pos_embed\"], [gs, gs],\n                                             num_prefix_tokens=bb.num_prefix_tokens)\nsd = {k: v for k, v in sd.items() if k not in tgt or tuple(v.shape) == tuple(tgt[k].shape)}\nmissing, _ = bb.load_state_dict(sd, strict=False)\nassert len(missing) < 0.1 * ntot, f\"only {ntot-len(missing)}/{ntot} keys loaded — do NOT train\"\nwith torch.no_grad():\n    _s = bb(torch.randn(2, 3, IMG_SIZE, IMG_SIZE))\n    _p = _s[:, bb.num_prefix_tokens:]\n    probe = torch.cat([_s[:, 0], _p.mean(1),\n                       _p.topk(max(1, _p.shape[1] // 8), dim=1).values.mean(1)], -1)\nassert probe.std() > 0.05, \"feature std ~0 — weights did not really load\"\nprint(f\"{layout} | loaded {ntot-len(missing)}/{ntot} keys | probe std {probe.std():.3f}\")\n\nclass PoolWrap(nn.Module):\n    def __init__(self, backbone, n_prefix):\n        super().__init__()\n        self.backbone, self.n_prefix = backbone, n_prefix\n    def forward(self, x):\n        t = self.backbone(x)\n        p = t[:, self.n_prefix:]\n        k = max(1, p.shape[1] // 8)\n        return torch.cat([t[:, 0], p.mean(1), p.topk(k, dim=1).values.mean(1)], -1)\n\nbb = bb.to(DEV).eval()\nfor p in bb.parameters():\n    p.requires_grad = False\nNB, NPRE = len(bb.blocks), bb.num_prefix_tokens\nMEAN = torch.tensor([0.485, 0.456, 0.406], device=DEV).view(1, 3, 1, 1)\nSTD  = torch.tensor([0.229, 0.224, 0.225], device=DEV).view(1, 3, 1, 1)\nprint(f\"{TIMM_NAME}: {NB} blocks, {sum(p.numel() for p in bb.parameters())/1e6:.0f}M params\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:09:25.252495Z","iopub.status.busy":"2026-09-19T06:09:25.252244Z","iopub.status.idle":"2026-09-19T06:09:27.482406Z","shell.execute_reply":"2026-09-19T06:09:27.481648Z"},"papermill":{"duration":2.237928,"end_time":"2026-09-19T06:09:27.484246+00:00","exception":false,"start_time":"2026-09-19T06:09:25.246318+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e3744ed3","cell_type":"markdown","source":"### 9 — Preprocessing and frozen features","metadata":{"papermill":{"duration":0.004825,"end_time":"2026-09-19T06:09:27.4939+00:00","exception":false,"start_time":"2026-09-19T06:09:27.489075+00:00","status":"completed"},"tags":[]}},{"id":"db67ab34","cell_type":"code","source":"INF_IMG_SZ = EFF_IMG if EFF_MODE else IMG_SIZE\nINF_K      = EFF_K   if EFF_MODE else K_PER_SLOT\n\nORDER_TAGS = [\"ImagePositionPatient\", \"ImageOrientationPatient\", \"InstanceNumber\"]\n\nclass StudyPrep(Dataset):\n    def __init__(self, ids, pick, lat, sz, ks):\n        self.ids, self.pick, self.lat, self.sz, self.ks = ids, pick, lat, sz, ks\n        self.ntk = sum(ks)\n    def __len__(self):\n        return len(self.ids)\n    def __getitem__(self, i):\n        sid = self.ids[i]\n        LT = self.lat.get(sid, \"R\")\n        vol = np.zeros((self.ntk, self.sz, self.sz), np.uint8)\n        mask = np.zeros(4, np.uint8)\n        for sl in range(4):\n            t0, K = int(np.sum(self.ks[:sl])), self.ks[sl]\n            if (sid, sl) not in self.pick:\n                continue\n            info = self.pick[(sid, sl)]\n            files = glob.glob(os.path.join(info[\"dir\"], DCM_GLOB))\n            if len(files) < 3:\n                continue\n            # Two-pass read. Filenames here are SOP UIDs and carry no spatial order, so\n            # sorting by any number in them scrambles the volume. Pass 1 reads geometry\n            # tags only and orders by ImagePositionPatient projected on the slice normal;\n            # pass 2 decodes just the slices kept.\n            # ~100 header reads per study, and each is network latency rather than CPU.\n            # Latency-bound work wants more concurrency than cores, so thread inside the worker.\n            with ThreadPoolExecutor(HDR_THREADS) as _ex:\n                hdr = list(_ex.map(lambda f: pydicom.dcmread(\n                    f, stop_before_pixels=True, force=True, specific_tags=ORDER_TAGS), files))\n            pos = []\n            for d in hdr:\n                iop = getattr(d, \"ImageOrientationPatient\", None)\n                ipp = getattr(d, \"ImagePositionPatient\", None)\n                if iop is not None and ipp is not None and len(iop) == 6 and len(ipp) == 3:\n                    nrm = np.cross(np.array(iop[:3], float), np.array(iop[3:], float))\n                    pos.append(float(np.dot(np.array(ipp, float), nrm)))\n                else:\n                    pos.append(None)\n            if all(p is not None for p in pos):\n                order = sorted(range(len(files)),\n                               key=lambda i: (pos[i], int(getattr(hdr[i], \"InstanceNumber\", 0) or 0)))\n            else:\n                order = sorted(range(len(files)),\n                               key=lambda i: int(getattr(hdr[i], \"InstanceNumber\", i) or i))\n            files = [files[i] for i in order]\n            n = len(files)\n            lo = int(0.10 * n); hi = max(int(0.90 * n) - 1, lo + 1)\n            idxs = np.linspace(lo, hi, K).round().astype(int)\n            if info[\"plane\"] == \"sagittal\" and LT == \"L\":\n                idxs = idxs[::-1]\n            got = 0\n            for kk, ii in enumerate(idxs):\n                try:\n                    d = pydicom.dcmread(files[int(ii)], force=True)\n                    a = d.pixel_array.astype(np.float32)\n                except Exception:\n                    continue\n                a = a * float(getattr(d, \"RescaleSlope\", 1) or 1) + \\\n                        float(getattr(d, \"RescaleIntercept\", 0) or 0)\n                lv_, hv_ = np.percentile(a, 0.5), np.percentile(a, 99.5)\n                a = np.clip((a - lv_) / max(hv_ - lv_, 1e-3), 0, 1)\n                sp = getattr(d, \"PixelSpacing\", [1.0, 1.0])\n                try:    spy, spx = float(sp[0]), float(sp[1])\n                except Exception: spy = spx = 1.0\n                hp = int(round(CROP_MM / max(spy, 1e-3))); wp = int(round(CROP_MM / max(spx, 1e-3)))\n                Hh, Ww = a.shape; cy, cx = Hh // 2, Ww // 2\n                c = a[max(0, cy - hp // 2):min(Hh, cy + hp // 2),\n                      max(0, cx - wp // 2):min(Ww, cx + wp // 2)]\n                if c.size == 0:\n                    continue\n                ph, pw = max(0, hp - c.shape[0]), max(0, wp - c.shape[1])\n                if ph or pw:\n                    c = np.pad(c, ((ph // 2, ph - ph // 2), (pw // 2, pw - pw // 2)))\n                img = cv2.resize(c, (self.sz, self.sz), interpolation=cv2.INTER_AREA)\n                if info[\"plane\"] in (\"axial\", \"coronal\") and LT == \"L\":\n                    img = np.ascontiguousarray(img[:, ::-1])\n                vol[t0 + kk] = (img * 255).astype(np.uint8)\n                got += 1\n            if got >= K // 2:\n                mask[sl] = 1\n        return torch.from_numpy(vol), torch.from_numpy(mask)\n\njobs = []\nif STAGE == \"train\":\n    jobs.append((\"train\", train_df[ID].tolist(), IMG_SIZE, K_PER_SLOT))\nif (meta.split == \"test\").any():\n    jobs.append((\"test\", test_df[ID].tolist(), INF_IMG_SZ, INF_K))\n\nstore = {}\nfor SPLIT, ids, SZ, KS in jobs:\n    NS_, NTK = len(ids), sum(KS)\n    fpath, mpath = f\"{WORK}/feat_{SPLIT}_{SZ}_{NTK}.npy\", f\"{WORK}/mask_{SPLIT}_{NTK}.npy\"\n    for _d in (CACHE_DIR, WORK):\n        if os.path.exists(f\"{_d}/feat_{SPLIT}_{SZ}_{NTK}.npy\"):\n            fpath, mpath = f\"{_d}/feat_{SPLIT}_{SZ}_{NTK}.npy\", f\"{_d}/mask_{SPLIT}_{NTK}.npy\"\n            break\n    need_px   = SPLIT == \"test\"\n    need_jpeg = SPLIT == \"train\" and SAVE_JPEG and not os.path.exists(f\"{CACHE_DIR}/jpg_train.bin\")\n    if os.path.exists(fpath) and not need_px:\n        store[SPLIT] = dict(feat=np.load(fpath), mask=np.load(mpath), ids=ids, ntok=NTK, sz=SZ)\n        print(f\"{SPLIT}: reusing {fpath}\"); continue\n\n    sub = meta[meta.split == SPLIT]\n    pick, lat = {}, {}\n    for _, r in sub[sub.slot >= 0].iterrows():\n        k = (r.sid, int(r.slot))\n        if k not in pick or r.n > pick[k][\"n\"]:\n            pick[k] = dict(dir=r.dir, n=int(r.n), plane=r.plane)\n    for sid, g in sub.groupby(\"sid\"):\n        lv = g.loc[g.L != \"\", \"L\"]\n        lat[sid] = lv.mode().iloc[0] if len(lv) else \"R\"\n\n    FEAT = np.zeros((NS_, NTK, EMB_DIM), np.float16)\n    MASK = np.zeros((NS_, 4), np.uint8)\n    PX   = np.zeros((NS_, NTK, SZ, SZ), np.uint8) if need_px else None\n    offs, cursor, done = np.zeros(NS_ * NTK + 1, np.int64), 0, 0\n    jf = open(f\"{WORK}/jpg_{SPLIT}.bin\", \"wb\") if need_jpeg else None\n    dp = nn.DataParallel(PoolWrap(bb, NPRE), device_ids=GPU_IDS)\n    dl = DataLoader(StudyPrep(ids, pick, lat, SZ, KS), batch_size=EXTRACT_BATCH, shuffle=False,\n                    num_workers=NUM_WORKERS, pin_memory=True, prefetch_factor=4)\n\n    for vols, masks in dl:\n        B = vols.shape[0]\n        MASK[done:done + B] = masks.numpy()\n        vnp = vols.numpy()\n        if jf is not None:\n            for b in range(B):\n                for t in range(NTK):\n                    offs[(done + b) * NTK + t] = cursor\n                    if vnp[b, t].max() > 0:\n                        okj, enc = cv2.imencode(\".jpg\", vnp[b, t],\n                                                [int(cv2.IMWRITE_JPEG_QUALITY), JPEG_Q])\n                        if okj:\n                            jf.write(enc.tobytes()); cursor += len(enc)\n                offs[(done + b) * NTK + NTK] = cursor\n        if PX is not None:\n            PX[done:done + B] = vnp\n        xb = vols.to(DEV, non_blocking=True).float() / 255.\n        pv = torch.cat([xb[:, :1], xb[:, :-1]], 1)\n        nx = torch.cat([xb[:, 1:], xb[:, -1:]], 1)\n        im = ((torch.stack([pv, xb, nx], 2).reshape(-1, 3, SZ, SZ)) - MEAN) / STD\n        with torch.inference_mode(), torch.cuda.amp.autocast():\n            f = dp(im)\n        FEAT[done:done + B] = f.view(B, NTK, EMB_DIM).half().cpu().numpy()\n        done += B\n        if done % (EXTRACT_BATCH * 50) < EXTRACT_BATCH:\n            print(f\"  {SPLIT} {done}/{NS_}  {(time.time()-T0)/60:.1f}m\"\n                  f\"{f'  jpeg {cursor/1e9:.2f}GB' if jf else ''}\")\n\n    if jf is not None:\n        jf.close(); np.save(f\"{WORK}/jpg_{SPLIT}_off.npy\", offs)\n    np.save(fpath, FEAT); np.save(mpath, MASK)\n    store[SPLIT] = dict(feat=FEAT, mask=MASK, ids=ids, ntok=NTK, sz=SZ, px=PX)\n    print(f\"{SPLIT}: {FEAT.shape} = {FEAT.nbytes/1e6:.0f} MB | slots {MASK.mean(0).round(2)}\")\n\nJOFF = JBLOB = None\nif \"train\" in store:\n    FEAT, MASK = store[\"train\"][\"feat\"], store[\"train\"][\"mask\"]\n    JDIR = next((d for d in (CACHE_DIR, WORK) if os.path.exists(f\"{d}/jpg_train_off.npy\")), None)\n    if JDIR:\n        JOFF = np.load(f\"{JDIR}/jpg_train_off.npy\")\n        JBLOB = np.memmap(f\"{JDIR}/jpg_train.bin\", np.uint8, \"r\")\n    print(f\"jpeg cache: {JDIR}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:09:27.504671Z","iopub.status.busy":"2026-09-19T06:09:27.504436Z","iopub.status.idle":"2026-09-19T06:34:54.24806Z","shell.execute_reply":"2026-09-19T06:34:54.246875Z"},"papermill":{"duration":1526.75148,"end_time":"2026-09-19T06:34:54.249793+00:00","exception":false,"start_time":"2026-09-19T06:09:27.498313+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d336eaf2","cell_type":"markdown","source":"### 10 — Head and Phase A","metadata":{"papermill":{"duration":0.006915,"end_time":"2026-09-19T06:34:54.261885+00:00","exception":false,"start_time":"2026-09-19T06:34:54.25497+00:00","status":"completed"},"tags":[]}},{"id":"4c4399f7","cell_type":"code","source":"SLOT_PRIOR_TABLE = {\"ACL\": (1, 3), \"MCL\": (2,), \"Medial Meniscus\": (1, 2, 3),\n                    \"Lateral Meniscus\": (1, 2, 3), \"Medial OA\": (2, 3), \"Lateral OA\": (2, 3),\n                    \"PF OA\": (0,), \"Effusion\": (0, 3), \"Synovitis\": (0, 3), \"Baker's\": (0, 3),\n                    \"Contusion\": (0, 2, 3), \"Fracture\": (0, 1, 2, 3)}\nSLOT_PRIOR = torch.zeros(NL, 4)\nfor _t, _sl in SLOT_PRIOR_TABLE.items():\n    SLOT_PRIOR[LABELS.index(_t), list(_sl)] = PRIOR_W\nSLOT_PRIOR = SLOT_PRIOR.to(DEV)\n\ndef head_logits(tk, mtok, H, tokmap, training=False):\n    h = H[\"pr\"](tk) + H[\"se\"][tokmap][None]\n    a = torch.einsum(\"bnh,oh->bon\", h, H[\"q\"]) / HEAD_HIDDEN ** 0.5\n    a = a + SLOT_PRIOR[:, tokmap][None]\n    a = torch.softmax(a.masked_fill(mtok.unsqueeze(1) < .5, -1e4), -1)\n    ctx = torch.einsum(\"bon,bnh->boh\", a, h)\n    if MAX_MIX > 0:\n        ctx = (1 - MAX_MIX) * ctx + MAX_MIX * torch.gather(\n            h, 1, a.argmax(-1).unsqueeze(-1).expand(-1, -1, HEAD_HIDDEN))\n    if training:\n        ctx = F.dropout(ctx, HEAD_DROP, True)\n    lg = (ctx * H[\"ow\"][None]).sum(-1) + H[\"ob\"]\n    return lg + LAB_MIX * (torch.softmax(H[\"mx\"], -1) @ lg.unsqueeze(-1)).squeeze(-1), h\n\nTK = torch.from_numpy(TOK2SLOT).to(DEV)\nNS = len(train_df)\noofA = np.full((NS, NL), np.nan, np.float32)\noofB = np.full((NS, NL), np.nan, np.float32)\nheads = {}\n\nif STAGE == \"infer\":\n    heads = torch.load(f\"{CKPT_DIR}/heads.pt\", map_location=\"cpu\")[\"heads\"]\n    print(f\"infer mode: {len(heads)} heads from {CKPT_DIR}\")\nelse:\n    Fg = torch.from_numpy(FEAT).float().to(DEV)\n    Yg, Wg = torch.from_numpy(soft).to(DEV), torch.from_numpy(sw).to(DEV)\n    Mg, Tg = torch.from_numpy(MASK.astype(np.float32)).to(DEV), torch.from_numpy(txt_t).to(DEV)\n\nfor FOLD in (FOLDS if STAGE == \"train\" else []):\n    tr = torch.from_numpy(np.where(train_df.fold.values != FOLD)[0]).to(DEV)\n    va = torch.from_numpy(np.where(train_df.fold.values == FOLD)[0]).to(DEV)\n    H = {\"pr\": nn.Sequential(nn.LayerNorm(EMB_DIM), nn.Linear(EMB_DIM, HEAD_HIDDEN),\n                             nn.GELU()).to(DEV),\n         \"se\": nn.Parameter(torch.randn(4, HEAD_HIDDEN, device=DEV) * .02),\n         \"q\":  nn.Parameter(torch.randn(NL, HEAD_HIDDEN, device=DEV) * .02),\n         \"ow\": nn.Parameter(torch.randn(NL, HEAD_HIDDEN, device=DEV) * .02),\n         \"ob\": nn.Parameter(torch.zeros(NL, device=DEV)),\n         \"mx\": nn.Parameter(torch.eye(NL, device=DEV) * 3.)}\n    th = nn.Linear(HEAD_HIDDEN, 64).to(DEV)\n    hp = [p for v in H.values() for p in (v.parameters() if hasattr(v, \"parameters\") else [v])] \\\n         + list(th.parameters())\n    opt = torch.optim.AdamW(hp, lr=LR_HEAD, weight_decay=WD)\n    sc = torch.optim.lr_scheduler.CosineAnnealingLR(opt, HEAD_EPOCHS)\n    best, bstate = -1., None\n\n    for ep in range(HEAD_EPOCHS):\n        perm = tr[torch.randperm(len(tr), device=DEV)]\n        for b0 in range(0, len(perm), 256):\n            sel = perm[b0:b0 + 256]\n            mtok = Mg[sel][:, TK]\n            lg, tk = head_logits(Fg[sel], mtok, H, TK, True)\n            gpool = (tk * mtok.unsqueeze(-1)).sum(1) / mtok.sum(1, keepdim=True).clamp(min=1)\n            loss = (F.binary_cross_entropy_with_logits(lg, Yg[sel], reduction=\"none\") * Wg[sel]).mean() \\\n                 + TXT_W * (1. - (F.normalize(th(gpool), dim=-1) * Tg[sel]).sum(-1)).mean()\n            opt.zero_grad(set_to_none=True); loss.backward()\n            torch.nn.utils.clip_grad_norm_(hp, 3.); opt.step()\n        sc.step()\n        if ep % 5 == 4 or ep == HEAD_EPOCHS - 1:\n            with torch.no_grad():\n                vp = torch.sigmoid(head_logits(Fg[va], Mg[va][:, TK], H, TK)[0]).cpu().numpy()\n            vi = va.cpu().numpy()\n            au = [roc_auc_score((soft[vi, j] > .5) * 1., vp[:, j]) for j in range(NL)\n                  if 0 < (soft[vi, j] > .5).sum() < len(vi)]\n            _take = (np.mean(au) > best) if PHASE_A_SELECT == \"val\" else (ep == HEAD_EPOCHS - 1)\n            if _take:\n                best, oofA[vi] = np.mean(au), vp\n                bstate = {k: (v.state_dict() if hasattr(v, \"state_dict\") else v.detach().clone())\n                          for k, v in H.items()}\n    heads[FOLD] = bstate\n    std = oofA[va.cpu().numpy()].std()\n    print(f\"fold {FOLD}: soft {best:.4f} | pred_std {std:.4f}\"\n          + (\"   !! COLLAPSED\" if std < .02 else \"\"))\n\nif STAGE == \"train\":\n    torch.save({\"heads\": heads}, f\"{WORK}/heads.pt\")\n    np.save(f\"{WORK}/oofA.npy\", oofA)\n    dm = ~np.isnan(oofA[:, 0]); gm = dm & (train_df.is_gold.values == 1)\n    ag = [roc_auc_score(train_df.loc[gm, L].values.astype(float), oofA[gm, j])\n          for j, L in enumerate(LABELS) if 0 < train_df.loc[gm, L].sum() < gm.sum()]\n    print(f\"Phase A OOF: gold n={int(gm.sum())} {np.mean(ag):.4f} -> LB ~{np.mean(ag)+0.042:.3f}\"\n          f\" | {(time.time()-T0)/60:.0f} min\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:34:54.27509Z","iopub.status.busy":"2026-09-19T06:34:54.274427Z","iopub.status.idle":"2026-09-19T06:36:19.277995Z","shell.execute_reply":"2026-09-19T06:36:19.276987Z"},"papermill":{"duration":85.012829,"end_time":"2026-09-19T06:36:19.279973+00:00","exception":false,"start_time":"2026-09-19T06:34:54.267144+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"38c4ee68","cell_type":"markdown","source":"### 11 — Phase B","metadata":{"papermill":{"duration":0.005254,"end_time":"2026-09-19T06:36:19.291047+00:00","exception":false,"start_time":"2026-09-19T06:36:19.285793+00:00","status":"completed"},"tags":[]}},{"id":"326b3771","cell_type":"code","source":"class JpegStudies(Dataset):\n    def __init__(self, idx):\n        self.idx = np.asarray(idx)\n    def __len__(self):\n        return len(self.idx)\n    def __getitem__(self, i):\n        s = int(self.idx[i])\n        vol = np.zeros((N_TOK, IMG_SIZE, IMG_SIZE), np.uint8)\n        for t in range(N_TOK):\n            o0, o1 = JOFF[s * N_TOK + t], JOFF[s * N_TOK + t + 1]\n            if o1 > o0:\n                vol[t] = cv2.imdecode(np.asarray(JBLOB[o0:o1]), cv2.IMREAD_GRAYSCALE)\n        return torch.from_numpy(vol), s\n\nif STAGE == \"train\":\n    del Fg, Tg                       # 0.62 GB of float32 features Phase B never touches\n    torch.cuda.empty_cache()\n\nINIT_SD = {k: v.detach().clone() for k, v in bb.state_dict().items()}\n\nfor FOLD in (FOLDS if STAGE == \"train\" and RUN_PHASE_B and JOFF is not None else []):\n    if os.path.exists(f\"{WORK}/m_f{FOLD}.pt\") or os.path.exists(f\"{CKPT_DIR}/m_f{FOLD}.pt\"):\n        print(f\"fold {FOLD}: checkpoint exists, skipping\"); continue\n    if time.time() - T0 > TIME_BUDGET - 1.2 * FOLD_SECONDS:\n        print(f\"fold {FOLD}: not enough time left\"); break\n    tri, vai = np.where(train_df.fold.values != FOLD)[0], np.where(train_df.fold.values == FOLD)[0]\n    bb.load_state_dict(INIT_SD)\n    for p in bb.parameters():\n        p.requires_grad = False\n    for bi in range(NB - UNFREEZE, NB):\n        for p in bb.blocks[bi].parameters():\n            p.requires_grad = True\n    for p in bb.norm.parameters():\n        p.requires_grad = True\n    bp = [p for p in bb.parameters() if p.requires_grad]\n\n    st = heads[FOLD]\n    H = {\"pr\": nn.Sequential(nn.LayerNorm(EMB_DIM), nn.Linear(EMB_DIM, HEAD_HIDDEN),\n                             nn.GELU()).to(DEV),\n         \"se\": nn.Parameter(st[\"se\"].clone().to(DEV)), \"q\": nn.Parameter(st[\"q\"].clone().to(DEV)),\n         \"ow\": nn.Parameter(st[\"ow\"].clone().to(DEV)), \"ob\": nn.Parameter(st[\"ob\"].clone().to(DEV)),\n         \"mx\": nn.Parameter(st[\"mx\"].clone().to(DEV))}\n    H[\"pr\"].load_state_dict(st[\"pr\"])\n    hp = [p for v in H.values() for p in (v.parameters() if hasattr(v, \"parameters\") else [v])]\n\n    grp = [{\"params\": list(bb.blocks[bi].parameters()), \"lr\": LR_BB * (0.75 ** (NB - 1 - bi))}\n           for bi in range(NB - UNFREEZE, NB)]\n    grp += [{\"params\": list(bb.norm.parameters()), \"lr\": LR_BB}, {\"params\": hp, \"lr\": LR_HEAD * .1}]\n    opt = torch.optim.AdamW(grp, weight_decay=WD)\n    tl = DataLoader(JpegStudies(tri), batch_size=BATCH_STUDIES, shuffle=True, drop_last=True,\n                    num_workers=NUM_WORKERS, pin_memory=True, persistent_workers=True,\n                    prefetch_factor=4)\n    vl = DataLoader(JpegStudies(vai), batch_size=BATCH_STUDIES * 2, shuffle=False,\n                    num_workers=NUM_WORKERS, pin_memory=True)\n    sch = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=[g[\"lr\"] for g in grp],\n                                              total_steps=max(1, PHASE_B_EPOCHS * len(tl) // ACCUM),\n                                              pct_start=.15, div_factor=10., final_div_factor=100.)\n    scaler = torch.cuda.amp.GradScaler()\n    ema = ModelEmaV3(bb, decay=EMA_DECAY, use_warmup=True, device=DEV)\n    dp = nn.DataParallel(PoolWrap(bb, NPRE), device_ids=GPU_IDS)\n\n    for ep in range(PHASE_B_EPOCHS):\n        bb.train(); opt.zero_grad(set_to_none=True)\n        for it, (xb_u8, sel) in enumerate(tl):\n            sel = sel.to(DEV)\n            xb = xb_u8.to(DEV, non_blocking=True).float() / 255.\n            B = xb.shape[0]; mb = Mg[sel].clone()\n            xb = xb.clamp(1e-4, 1).pow(torch.empty(B, 1, 1, 1, device=DEV).uniform_(.8, 1.25))\n            xb = xb + torch.randn_like(xb) * .01\n            sh = int(IMG_SIZE * .04)\n            xb = torch.roll(xb, (np.random.randint(-sh, sh+1), np.random.randint(-sh, sh+1)), (2, 3))\n            mb = mb * (torch.rand(B, 4, device=DEV) > .10).float()\n            mb[:, 0] = torch.clamp(mb[:, 0] + (mb.sum(1) == 0).float(), max=1.)\n            pv, nx = torch.cat([xb[:, :1], xb[:, :-1]], 1), torch.cat([xb[:, 1:], xb[:, -1:]], 1)\n            im = ((torch.stack([pv, xb, nx], 2).reshape(-1, 3, IMG_SIZE, IMG_SIZE)) - MEAN) / STD\n            with torch.cuda.amp.autocast():\n                tk0 = dp(im).view(B, N_TOK, EMB_DIM)\n                lg, _ = head_logits(tk0, mb[:, TK], H, TK)\n                loss = (F.binary_cross_entropy_with_logits(lg, Yg[sel], reduction=\"none\")\n                        * Wg[sel]).mean()\n            scaler.scale(loss / ACCUM).backward()\n            if (it + 1) % ACCUM == 0:\n                scaler.unscale_(opt); torch.nn.utils.clip_grad_norm_(bp + hp, 3.)\n                scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True)\n                ema.update(bb)\n                if sch.last_epoch < sch.total_steps - 1:\n                    sch.step()\n            if it % 200 == 0 or (ep == 0 and it == 0):\n                print(f\"  f{FOLD} e{ep} {it}/{len(tl)} loss {loss.item():.4f} (floor {FLOOR:.3f}) \"\n                      f\"{torch.cuda.max_memory_allocated()/1e9:.1f}GB {(time.time()-T0)/60:.0f}m\")\n            if time.time() - T0 > TIME_BUDGET:\n                break\n        if time.time() - T0 > TIME_BUDGET:\n            print(\"  time budget hit\"); break\n\n    ema.module.eval()\n    dpe = nn.DataParallel(PoolWrap(ema.module, NPRE), device_ids=GPU_IDS)\n    vp = np.zeros((len(vai), NL), np.float32); off = 0\n    with torch.inference_mode():\n        for xb_u8, sel in vl:\n            xb = xb_u8.to(DEV).float() / 255.\n            B = xb.shape[0]\n            pv, nx = torch.cat([xb[:, :1], xb[:, :-1]], 1), torch.cat([xb[:, 1:], xb[:, -1:]], 1)\n            im = ((torch.stack([pv, xb, nx], 2).reshape(-1, 3, IMG_SIZE, IMG_SIZE)) - MEAN) / STD\n            with torch.cuda.amp.autocast():\n                tk0 = dpe(im).view(B, N_TOK, EMB_DIM)\n                lg, _ = head_logits(tk0, Mg[sel.to(DEV)][:, TK], H, TK)\n            vp[off:off + B] = torch.sigmoid(lg.float()).cpu().numpy(); off += B\n    oofB[vai] = vp\n    af = [roc_auc_score((soft[vai, j] > .5) * 1., vp[:, j]) for j in range(NL)\n          if 0 < (soft[vai, j] > .5).sum() < len(vai)]\n    print(f\"  fold {FOLD} soft {np.mean(af):.4f} | {(time.time()-T0)/60:.0f} min\")\n    torch.save({\"bb\": {k: v.half() for k, v in ema.module.state_dict().items()},\n                **{k: (v.state_dict() if hasattr(v, \"state_dict\") else v.detach().cpu())\n                   for k, v in H.items()}}, f\"{WORK}/m_f{FOLD}.pt\")\nif STAGE == \"train\":\n    np.save(f\"{WORK}/oofB.npy\", oofB)\n    if not KEEP_JPEG and not np.isnan(oofB[:, 0]).all():\n        del JBLOB\n        for p in (f\"{WORK}/jpg_train.bin\", f\"{WORK}/jpg_train_off.npy\"):\n            if os.path.exists(p):\n                os.remove(p)\n        print(\"removed the JPEG cache (KEEP_JPEG=False)\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T06:36:19.303524Z","iopub.status.busy":"2026-09-19T06:36:19.303022Z","iopub.status.idle":"2026-09-19T15:05:36.418502Z","shell.execute_reply":"2026-09-19T15:05:36.417378Z"},"papermill":{"duration":30557.132456,"end_time":"2026-09-19T15:05:36.42895+00:00","exception":false,"start_time":"2026-09-19T06:36:19.296494+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"06454615","cell_type":"markdown","source":"### 12 — Out-of-fold evaluation","metadata":{"papermill":{"duration":0.011223,"end_time":"2026-09-19T15:05:36.448755+00:00","exception":false,"start_time":"2026-09-19T15:05:36.437532+00:00","status":"completed"},"tags":[]}},{"id":"7da0c85e","cell_type":"code","source":"for nm, oo in [(\"Phase A\", oofA), (\"Phase B\", oofB)]:\n    if np.isnan(oo[:, 0]).all():\n        continue\n    dm = ~np.isnan(oo[:, 0]); gm = dm & (train_df.is_gold.values == 1)\n    ag = [roc_auc_score(train_df.loc[gm, L].values.astype(float), oo[gm, j])\n          for j, L in enumerate(LABELS) if 0 < train_df.loc[gm, L].sum() < gm.sum()]\n    asf = [roc_auc_score((soft[dm, j] > .5) * 1., oo[dm, j]) for j in range(NL)\n           if 0 < (soft[dm, j] > .5).sum() < dm.sum()]\n    print(f\"{nm:8s} gold n={int(gm.sum())}: {np.mean(ag):.4f}  soft n={int(dm.sum())}: {np.mean(asf):.4f}\"\n          f\"   -> LB ~{np.mean(ag)+0.042:.3f}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T15:05:36.467047Z","iopub.status.busy":"2026-09-19T15:05:36.466587Z","iopub.status.idle":"2026-09-19T15:05:36.586739Z","shell.execute_reply":"2026-09-19T15:05:36.585805Z"},"papermill":{"duration":0.131479,"end_time":"2026-09-19T15:05:36.588343+00:00","exception":false,"start_time":"2026-09-19T15:05:36.456864+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bc929847-4a66-478d-9034-7a30e28d650a","cell_type":"markdown","source":"### 12b — Per-finding report (v2)","metadata":{}},{"id":"2ecbf64c-3efa-4bde-81ff-854ace416c78","cell_type":"code","source":"import json\nfrom sklearn.metrics import roc_auc_score\n\ndef _auc(y, p):\n    return roc_auc_score(y, p) if 0 < y.sum() < len(y) else np.nan\n\ndef report_auc(name, oo, n_boot=N_BOOT, seed=0):\n    dm = ~np.isnan(oo[:, 0])\n    if not dm.any():\n        return None\n    gm = dm & (train_df.is_gold.values == 1)\n    yg_ = train_df.loc[gm, LABELS].values.astype(float)\n    rows = []\n    for j, L in enumerate(LABELS):\n        rows.append((L, _auc(yg_[:, j], oo[gm, j]), _auc((soft[dm, j] > .5) * 1., oo[dm, j]),\n                     int(yg_[:, j].sum())))\n    rng, boots = np.random.RandomState(seed), []\n    for _ in range(n_boot):\n        b = rng.randint(0, len(yg_), len(yg_))\n        v = [_auc(yg_[b, j], oo[gm, j][b]) for j in range(NL)]\n        v = [x for x in v if not np.isnan(x)]\n        if v:\n            boots.append(np.mean(v))\n    lo, hi = np.percentile(boots, [2.5, 97.5]) if boots else (np.nan, np.nan)\n    print(f\"\\n{name}: per-finding AUC   (gold n={int(gm.sum())} | soft n={int(dm.sum())})\")\n    print(f\"  {'finding':18s} {'gold':>7s} {'soft':>7s} {'gold pos':>9s}\")\n    for L, ag_, as_, npos in rows:\n        print(f\"  {L:18s} {ag_:7.3f} {as_:7.3f} {npos:9d}\")\n    mg = np.nanmean([r[1] for r in rows]); ms = np.nanmean([r[2] for r in rows])\n    print(f\"  {'mean':18s} {mg:7.3f} {ms:7.3f}   gold 95% bootstrap CI [{lo:.3f}, {hi:.3f}]\")\n    return dict(gold=float(mg), soft=float(ms), gold_ci=[float(lo), float(hi)])\n\n_info = dict(preset=PRESET, backbone=TIMM_NAME, img=IMG_SIZE, seed=SEED, phase_b=RUN_PHASE_B,\n             label_source=LABEL_SOURCE, phase_a_select=PHASE_A_SELECT, folds_from=FOLD_SRC)\nfor _nm, _oo in [(\"Phase A\", oofA), (\"Phase B\", oofB)]:\n    _r = report_auc(_nm, _oo)\n    if _r:\n        _info[_nm.replace(\" \", \"_\").lower()] = _r\njson.dump(_info, open(f\"{WORK}/run_info.json\", \"w\"), indent=2)\nprint(\"\\nrun_info.json written. Wide gold intervals mean the gold number cannot separate models that are close.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"92274668","cell_type":"markdown","source":"### 13 — Inference","metadata":{"papermill":{"duration":0.008767,"end_time":"2026-09-19T15:05:36.60564+00:00","exception":false,"start_time":"2026-09-19T15:05:36.596873+00:00","status":"completed"},"tags":[]}},{"id":"2204d9f2","cell_type":"code","source":"assert \"test\" in store, \"no test split — Kaggle mounts the real test set only at scoring time\"\ntst = store[\"test\"]\nNT, NTK_T, SZ_T = len(tst[\"ids\"]), tst[\"ntok\"], tst[\"sz\"]\nTKT = torch.from_numpy(np.concatenate([np.full(k, i, dtype=np.int64)\n                                       for i, k in enumerate(INF_K)])).to(DEV)\nMTt = torch.from_numpy(tst[\"mask\"].astype(np.float32)).to(DEV)\nMTt[:, 0] = torch.clamp(MTt[:, 0] + (MTt.sum(1) == 0).float(), max=1.)\nck = sorted(glob.glob(f\"{CKPT_DIR}/m_f*.pt\") + glob.glob(f\"{WORK}/m_f*.pt\"))\nck = sorted(set(ck))[:1 if EFF_MODE else None]\nrank, used = np.zeros((NT, NL)), 0\nprint(f\"{len(ck)} checkpoints | {NT} test studies | {NTK_T} tokens @ {SZ_T}px\")\n\nif not ck:\n    Ft = torch.from_numpy(tst[\"feat\"]).float().to(DEV)\n    for st in heads.values():\n        H = {\"pr\": nn.Sequential(nn.LayerNorm(EMB_DIM), nn.Linear(EMB_DIM, HEAD_HIDDEN),\n                                 nn.GELU()).to(DEV),\n             \"se\": st[\"se\"].to(DEV), \"q\": st[\"q\"].to(DEV), \"ow\": st[\"ow\"].to(DEV),\n             \"ob\": st[\"ob\"].to(DEV), \"mx\": st[\"mx\"].to(DEV)}\n        H[\"pr\"].load_state_dict(st[\"pr\"])\n        with torch.no_grad():\n            P = torch.sigmoid(head_logits(Ft, MTt[:, TKT], H, TKT)[0]).cpu().numpy()\n        for j in range(NL):\n            rank[:, j] += pd.Series(P[:, j]).rank(pct=True).values\n        used += 1\nelse:\n    if SZ_T != IMG_SIZE:\n        bb = timm.create_model(TIMM_NAME, pretrained=False, num_classes=0,\n                               img_size=SZ_T, global_pool=\"\").to(DEV)\n    for cp in ck:\n        st = torch.load(cp, map_location=\"cpu\")\n        bsd = {k: v.float() for k, v in st[\"bb\"].items()}   # checkpoints are stored fp16\n        if SZ_T != IMG_SIZE and \"pos_embed\" in bsd:\n            gs = SZ_T // bb.patch_embed.patch_size[0]\n            bsd[\"pos_embed\"] = resample_abs_pos_embed(bsd[\"pos_embed\"], [gs, gs],\n                                                      num_prefix_tokens=bb.num_prefix_tokens)\n        miss, _ = bb.load_state_dict(bsd, strict=False)\n        assert len(miss) < 5, f\"checkpoint failed to load: {miss[:5]}\"\n        H = {\"pr\": nn.Sequential(nn.LayerNorm(EMB_DIM), nn.Linear(EMB_DIM, HEAD_HIDDEN),\n                                 nn.GELU()).to(DEV),\n             \"se\": st[\"se\"].to(DEV), \"q\": st[\"q\"].to(DEV), \"ow\": st[\"ow\"].to(DEV),\n             \"ob\": st[\"ob\"].to(DEV), \"mx\": st[\"mx\"].to(DEV)}\n        H[\"pr\"].load_state_dict(st[\"pr\"])\n        bb.eval()\n        dp = nn.DataParallel(PoolWrap(bb, NPRE), device_ids=GPU_IDS)\n        P = np.zeros((NT, NL), np.float32)\n        for view in range(TTA):\n            with torch.inference_mode():\n                for b0 in range(0, NT, BATCH_STUDIES * 2):\n                    xb = torch.from_numpy(tst[\"px\"][b0:b0 + BATCH_STUDIES * 2]).to(DEV).float() / 255.\n                    if view:\n                        s_ = max(1, int(SZ_T * .03)) * (1 if view == 1 else -1)\n                        xb = torch.roll(xb, (s_, -s_), (2, 3))\n                    B = xb.shape[0]\n                    pv, nx = torch.cat([xb[:, :1], xb[:, :-1]], 1), torch.cat([xb[:, 1:], xb[:, -1:]], 1)\n                    im = ((torch.stack([pv, xb, nx], 2).reshape(-1, 3, SZ_T, SZ_T))\n                          - MEAN) / STD\n                    with torch.cuda.amp.autocast():\n                        tk0 = dp(im).view(B, NTK_T, EMB_DIM)\n                        lg, _ = head_logits(tk0, MTt[b0:b0 + B][:, TKT], H, TKT)\n                    P[b0:b0 + B] += torch.sigmoid(lg.float()).cpu().numpy() / TTA\n        for j in range(NL):\n            rank[:, j] += pd.Series(P[:, j]).rank(pct=True).values\n        used += 1\n        print(f\"  {os.path.basename(cp)} done  {(time.time()-T0)/60:.0f}m\")\n\nsub = pd.DataFrame({ID: tst[\"ids\"]})\nfor j, L in enumerate(LABELS):\n    sub[L] = rank[:, j] / max(used, 1)\nsub[list(sub_df.columns)].to_csv(\"submission.csv\", index=False)\nrt = time.time() - T0\nprint(f\"submission.csv {sub.shape} | {used} models | {rt/60:.1f} min \"\n      f\"| efficiency -2.22*AUC + {rt/32400:.4f}\")","metadata":{"execution":{"iopub.execute_input":"2026-09-19T15:05:36.624079Z","iopub.status.busy":"2026-09-19T15:05:36.623507Z","iopub.status.idle":"2026-09-19T15:05:44.801588Z","shell.execute_reply":"2026-09-19T15:05:44.800693Z"},"papermill":{"duration":8.189294,"end_time":"2026-09-19T15:05:44.803321+00:00","exception":false,"start_time":"2026-09-19T15:05:36.614027+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}