{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"NvidiaTeslaT4","isGpuEnabled":true,"isInternetEnabled":false,"sourceType":"notebook"}},"nbformat_minor":5,"nbformat":4,"cells":[{"id":"b7867db8-6cc7-41fe-acb4-6357bade186c","cell_type":"markdown","source":"# RSNA Knee Abnormality Detection — Kaggle notebook\n\nEnd-to-end pipeline:\nDICOM -> series selection -> slice preprocessing -> attention MIL -> 12 probabilities -> submission.csv\n\nThe notebook is designed for Kaggle GPU and uses the fixed competition input directory, DICOM spatial sorting, fluid-sensitive series selection, uint8 cache, FP16, smoke test, study-level GroupKFold, checkpoints, resume and final submission validation.\n\nRun cells from top to bottom. Enable GPU before running.","metadata":{}},{"id":"fd6a6794-ba76-4162-9d69-44d36ea04591","cell_type":"markdown","source":"## 0. Configuration\n\nThe default ResNet18 is the safest offline baseline. After the pipeline works, BACKBONE can be changed to dinov2 if the required timm model weights are available.","metadata":{}},{"id":"92d4578c-34f4-4bf2-a148-be8bf15eebe1","cell_type":"code","source":"from pathlib import Path\nimport os, gc, pickle, random, re, time, warnings\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport torch\n\nSEED = 42\nINPUT_BASE = Path(\"/kaggle/input\")\nWORK_DIR = Path(\"/kaggle/working/rsna_knee_abnormality_v1\")\nCACHE_DIR = WORK_DIR / \"cache\"\nCKPT_DIR = WORK_DIR / \"checkpoints\"\nfor p in [WORK_DIR, CACHE_DIR, CKPT_DIR]:\n    p.mkdir(parents=True, exist_ok=True)\n\nDATA_ROOT = Path(\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n)\nBACKBONE = \"dinov2\"     # optional: \"dinov2\"\nPRETRAINED = True\nATTENTION_IMPLEMENTATION = \"eager\"\nIMG_SIZE = 224\nMAX_SLICES = 32\nK_SLICES = 12\nN_SLOTS = 3\n\nN_FOLDS = 5\nFOLDS_TO_RUN = list(range(N_FOLDS))\nEPOCHS = 8\nPATIENCE = 3\nBATCH_SIZE = 4\nACCUM_STEPS = 2\nNUM_WORKERS = 2\nLEARNING_RATE = 2e-4\nWEIGHT_DECAY = 1e-4\nLABEL_SMOOTHING = 0.02\nGRAD_CLIP = 1.0\nUSE_AMP = True\nUSE_DATA_PARALLEL = True\n\nPREPARE_CACHE = True\nREBUILD_CACHE = False\nREBUILD_INDEX = False\nMAX_CACHE_STUDIES = None\nRUN_SMOKE_TEST = True\nRUN_FULL_TRAIN = True\nRUN_TEST_INFERENCE = True\nRESUME = True\nUSE_TTA = False  # intensity TTA only; no horizontal flip\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nif DEVICE != \"cuda\":\n    raise RuntimeError(\"Enable a Kaggle GPU accelerator before running this notebook.\")\ntorch.backends.cudnn.benchmark = True\ntorch.backends.cuda.matmul.allow_tf32 = True\ntry:\n    torch.set_float32_matmul_precision(\"high\")\nexcept Exception:\n    pass\n\nprint({\"device\": DEVICE, \"gpu_count\": torch.cuda.device_count(), \"backbone\": BACKBONE, \"attention\": ATTENTION_IMPLEMENTATION, \"folds\": FOLDS_TO_RUN})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:04:58.349059Z","iopub.execute_input":"2026-09-08T15:04:58.349296Z","iopub.status.idle":"2026-09-08T15:05:04.045476Z","shell.execute_reply.started":"2026-09-08T15:04:58.349274Z","shell.execute_reply":"2026-09-08T15:05:04.044628Z"}},"outputs":[{"name":"stdout","text":"{'device': 'cuda', 'gpu_count': 2, 'backbone': 'dinov2', 'attention': 'eager', 'folds': [0, 1, 2, 3, 4]}\n","output_type":"stream"}],"execution_count":1},{"id":"5227a96b-264c-411f-9286-cbfa16af5be8","cell_type":"markdown","source":"## 1. Imports, seed and fixed Kaggle data paths","metadata":{}},{"id":"14fdba35-da1a-4431-a63b-74cf27494b99","cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm.auto import tqdm\nwarnings.filterwarnings(\"ignore\")\n\ndef seed_everything(seed=SEED):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nseed_everything()\n\nROOT = DATA_ROOT\nTRAIN_CSV = ROOT / \"train.csv\"\nTEST_CSV = ROOT / \"test.csv\"\nSAMPLE_CSV = ROOT / \"sample_submission.csv\"\nfor required_path in [ROOT, TRAIN_CSV, TEST_CSV, SAMPLE_CSV]:\n    if not required_path.exists():\n        raise FileNotFoundError(f\"Missing Kaggle path: {required_path}\")\n\ntrain_raw = pd.read_csv(TRAIN_CSV)\ntest_raw = pd.read_csv(TEST_CSV)\nsample_raw = pd.read_csv(SAMPLE_CSV)\nprint(\"ROOT:\", ROOT)\nprint(\"train:\", train_raw.shape, \"test:\", test_raw.shape, \"sample:\", sample_raw.shape)\ndisplay(train_raw.head(2))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:05:04.04657Z","iopub.execute_input":"2026-09-08T15:05:04.047038Z","iopub.status.idle":"2026-09-08T15:05:06.008643Z","shell.execute_reply.started":"2026-09-08T15:05:04.04701Z","shell.execute_reply":"2026-09-08T15:05:06.008011Z"}},"outputs":[{"name":"stdout","text":"ROOT: /kaggle/input/competitions/rsna-knee-abnormality-detection\ntrain: (4407, 14) test: (3, 1) sample: (3, 13)\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"                                    StudyInstanceUID  \\\n0  1.2.826.0.1.3680043.8.498.10004873229099053869...   \n1  1.2.826.0.1.3680043.8.498.10004945927472656027...   \n\n                                              Report  ACL  MCL  \\\n0  Técnica: RMN de la rodilla. Resultados: Rotura...  NaN  NaN   \n1  [DATE]: * MR Knie Rechts 15ch AA Klinische Inl...  NaN  NaN   \n\n   Medial Meniscus  Lateral Meniscus  Medial OA  Lateral OA  PF OA  Effusion  \\\n0              NaN               NaN        NaN         NaN    NaN       NaN   \n1              NaN               NaN        NaN         NaN    NaN       NaN   \n\n   Synovitis  Baker's  Contusion  Fracture  \n0        NaN      NaN        NaN       NaN  \n1        NaN      NaN        NaN       NaN  ","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>StudyInstanceUID</th>\n      <th>Report</th>\n      <th>ACL</th>\n      <th>MCL</th>\n      <th>Medial Meniscus</th>\n      <th>Lateral Meniscus</th>\n      <th>Medial OA</th>\n      <th>Lateral OA</th>\n      <th>PF OA</th>\n      <th>Effusion</th>\n      <th>Synovitis</th>\n      <th>Baker's</th>\n      <th>Contusion</th>\n      <th>Fracture</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>1.2.826.0.1.3680043.8.498.10004873229099053869...</td>\n      <td>Técnica: RMN de la rodilla. Resultados: Rotura...</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>1.2.826.0.1.3680043.8.498.10004945927472656027...</td>\n      <td>[DATE]: * MR Knie Rechts 15ch AA Klinische Inl...</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":2},{"id":"4f3677c8-234e-4ede-bad4-948a0f5288c0","cell_type":"markdown","source":"## 2. Normalize IDs and infer the 12 target columns","metadata":{}},{"id":"cabf7acc-e2be-4434-8887-64ae02d0f21d","cell_type":"code","source":"def pick_id_col(df):\n    preferred = [\"StudyInstanceUID\", \"study_instance_uid\", \"study_id\", \"study_uid\", \"StudyUID\", \"uid\", \"id\", \"row_id\"]\n    low = {str(c).lower(): c for c in df.columns}\n    for name in preferred:\n        if name.lower() in low:\n            return low[name.lower()]\n    scored = []\n    for c in df.columns:\n        name = str(c).lower()\n        score = 3 * int(\"study\" in name) + 2 * int(\"uid\" in name) + int(\"id\" in name)\n        score += 2 * int(df[c].nunique(dropna=True) == len(df))\n        scored.append((score, c))\n    return max(scored, key=lambda x: x[0])[1]\n\ndef binary_like(s):\n    x = pd.to_numeric(s, errors=\"coerce\").dropna()\n    return len(x) > 0 and set(np.unique(x)).issubset({0, 1})\n\ndef norm_uid(s):\n    return s.astype(str).str.strip()\n\ntrain_id = pick_id_col(train_raw)\nsample_id = pick_id_col(sample_raw) if sample_raw is not None else train_id\ntest_id = pick_id_col(test_raw) if test_raw is not None else sample_id\n\nsample_targets = [c for c in sample_raw.columns if c != sample_id] if sample_raw is not None else []\nshared = [c for c in sample_targets if c in train_raw.columns]\nbinary_targets = [c for c in train_raw.columns if c != train_id and binary_like(train_raw[c])]\nTARGET_COLUMNS = shared if len(shared) >= 8 else binary_targets\nif len(TARGET_COLUMNS) > 12:\n    TARGET_COLUMNS = TARGET_COLUMNS[:12]\nif len(TARGET_COLUMNS) != 12:\n    raise ValueError(f\"Expected 12 targets, found {len(TARGET_COLUMNS)}: {TARGET_COLUMNS}\")\n\ntrain_df = train_raw.rename(columns={train_id: \"StudyInstanceUID\"}).copy()\ntrain_df[\"StudyInstanceUID\"] = norm_uid(train_df[\"StudyInstanceUID\"])\n\n# Most rows have a Report but no manually supplied binary labels. Do not turn\n# missing labels into false negatives. Build soft weak labels from the report,\n# then override them with the available gold labels.\nREPORT_COLUMN = \"Report\" if \"Report\" in train_df.columns else None\nREPORT_WEIGHT = 0.35\nGOLD_WEIGHT = 8.0\nREPORT_TERMS = {\n    \"ACL\": [\"acl\", \"anterior cruciate\", \"cruciate ligament\", \"ligamento cruzado anterior\", \"vordere kreuzband\"],\n    \"MCL\": [\"mcl\", \"medial collateral\", \"ligamento colateral medial\", \"mediales kollateralband\"],\n    \"Medial Meniscus\": [\"medial meniscus\", \"menisque medial\", \"meniskus medialis\", \"menisco medial\"],\n    \"Lateral Meniscus\": [\"lateral meniscus\", \"menisque lateral\", \"meniskus lateralis\", \"menisco lateral\"],\n    \"Medial OA\": [\"medial osteoarthritis\", \"medial compartment\", \"medial joint space\", \"medial femorotibial\", \"mediale gonarthrose\"],\n    \"Lateral OA\": [\"lateral osteoarthritis\", \"lateral compartment\", \"lateral joint space\", \"lateral femorotibial\", \"laterale gonarthrose\"],\n    \"PF OA\": [\"patellofemoral\", \"patello-femoral\", \"femoropatellar\", \"patellare arthrose\"],\n    \"Effusion\": [\"effusion\", \"joint fluid\", \"joint effusion\", \"epanchement\", \"derrame\", \"erguss\", \"versamento\"],\n    \"Synovitis\": [\"synovitis\", \"synovial thickening\", \"synovial hypertrophy\", \"synovite\"],\n    \"Baker's\": [\"baker\", \"popliteal cyst\", \"cyste poplitee\", \"popliteal bursa\"],\n    \"Contusion\": [\"contusion\", \"bone bruise\", \"marrow edema\", \"osseous edema\", \"prellung\", \"contusao\"],\n    \"Fracture\": [\"fracture\", \"fraktur\", \"fractura\", \"frattura\", \"fracture line\"],\n}\nNEGATION_TERMS = [\n    \"no \", \"without \", \"negative for\", \"not seen\", \"not identified\", \"intact\", \"normal\", \"unremarkable\", \"preserved\",\n    \"kein \", \"keine \", \"nicht \", \"ohne \", \"sin \", \"sem \", \"aucun \", \"pas de \", \"nessun \", \"nessuna \", \"non \",\n]\n\ndef normalize_report(text):\n    import unicodedata\n    text = unicodedata.normalize(\"NFKD\", str(text).lower())\n    return \"\".join(ch for ch in text if not unicodedata.combining(ch))\n\ndef weak_labels_from_report(report):\n    text = normalize_report(report)\n    clauses = re.split(r\"[.;,():\\n]+\", text)\n    values, weights = [], []\n    for target in TARGET_COLUMNS:\n        terms = REPORT_TERMS.get(target, [target.lower()])\n        states = []\n        for clause in clauses:\n            if any(term in clause for term in terms):\n                states.append(0.0 if any(neg in clause for neg in NEGATION_TERMS) else 1.0)\n        if not states:\n            values.append(0.5)\n            weights.append(0.0)\n        else:\n            values.append(1.0 if any(v == 1.0 for v in states) else 0.0)\n            weights.append(REPORT_WEIGHT)\n    return values, weights\n\nweak_values, weak_weights = [], []\nreports = train_df[REPORT_COLUMN] if REPORT_COLUMN else pd.Series([\"\"] * len(train_df))\nfor report in reports:\n    v, w = weak_labels_from_report(report)\n    weak_values.append(v)\n    weak_weights.append(w)\nweak_values = np.asarray(weak_values, dtype=\"float32\")\nweak_weights = np.asarray(weak_weights, dtype=\"float32\")\nraw_targets = train_df[TARGET_COLUMNS].apply(pd.to_numeric, errors=\"coerce\")\ngold_mask = raw_targets.notna().to_numpy()\neffective_targets = weak_values.copy()\neffective_weights = weak_weights.copy()\ngold_values = raw_targets.fillna(0).to_numpy(dtype=\"float32\")\neffective_targets[gold_mask] = gold_values[gold_mask]\neffective_weights[gold_mask] = GOLD_WEIGHT\nfor j, c in enumerate(TARGET_COLUMNS):\n    train_df[c] = effective_targets[:, j]\n    train_df[f\"__weight__{c}\"] = effective_weights[:, j]\nTARGET_WEIGHT_COLUMNS = [f\"__weight__{c}\" for c in TARGET_COLUMNS]\n\nif test_raw is not None:\n    test_df = test_raw[[test_id]].rename(columns={test_id: \"StudyInstanceUID\"}).copy()\nelse:\n    if sample_raw is None:\n        raise FileNotFoundError(\"Need test.csv or sample_submission.csv.\")\n    test_df = sample_raw[[sample_id]].rename(columns={sample_id: \"StudyInstanceUID\"}).copy()\ntest_df[\"StudyInstanceUID\"] = norm_uid(test_df[\"StudyInstanceUID\"])\n\nprint(\"TARGET_COLUMNS:\", TARGET_COLUMNS)\nprint(\"train studies:\", len(train_df), \"test studies:\", len(test_df))\nprint({\"gold_labels\": int(gold_mask.sum()), \"report_labels\": int((weak_weights > 0).sum()), \"report_column\": REPORT_COLUMN})\ndisplay(train_df[TARGET_COLUMNS].mean().sort_values().to_frame(\"positive_rate\").T)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:05:06.010739Z","iopub.execute_input":"2026-09-08T15:05:06.011175Z","iopub.status.idle":"2026-09-08T15:05:07.856696Z","shell.execute_reply.started":"2026-09-08T15:05:06.011115Z","shell.execute_reply":"2026-09-08T15:05:07.855908Z"}},"outputs":[{"name":"stdout","text":"TARGET_COLUMNS: ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\ntrain studies: 4407 test studies: 3\n{'gold_labels': 696, 'report_labels': 15272, 'report_column': 'Report'}\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"               Fracture  Contusion  Synovitis       ACL  Lateral Meniscus  \\\npositive_rate  0.460517    0.50953   0.530633  0.534831          0.544702   \n\n                Baker's  Lateral OA     PF OA  Effusion  Medial OA       MCL  \\\npositive_rate  0.545269    0.574881  0.575902  0.584638   0.586113  0.588496   \n\n               Medial Meniscus  \npositive_rate         0.607102  ","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>Fracture</th>\n      <th>Contusion</th>\n      <th>Synovitis</th>\n      <th>ACL</th>\n      <th>Lateral Meniscus</th>\n      <th>Baker's</th>\n      <th>Lateral OA</th>\n      <th>PF OA</th>\n      <th>Effusion</th>\n      <th>Medial OA</th>\n      <th>MCL</th>\n      <th>Medial Meniscus</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>positive_rate</th>\n      <td>0.460517</td>\n      <td>0.50953</td>\n      <td>0.530633</td>\n      <td>0.534831</td>\n      <td>0.544702</td>\n      <td>0.545269</td>\n      <td>0.574881</td>\n      <td>0.575902</td>\n      <td>0.584638</td>\n      <td>0.586113</td>\n      <td>0.588496</td>\n      <td>0.607102</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":3},{"id":"1df599e5-7c9e-46d9-b7aa-8a37db882af6","cell_type":"markdown","source":"## 3. Index DICOM and read series headers","metadata":{}},{"id":"0ce85d16-f6cd-407b-a90d-4050ee1f9e27","cell_type":"code","source":"if not (ROOT / \"train_series\").exists() or not (ROOT / \"test_series\").exists():\n    raise FileNotFoundError(\"Expected train_series and test_series under the competition input directory.\")\nDICOM_ROOTS = [ROOT]\nALL_IDS = set(train_df[\"StudyInstanceUID\"]) | set(test_df[\"StudyInstanceUID\"])\nINDEX_PATH = WORK_DIR / \"series_index.pkl\"\n\ndef build_series_index(roots, study_ids, rebuild=False):\n    if INDEX_PATH.exists() and not rebuild:\n        with open(INDEX_PATH, \"rb\") as f:\n            return pickle.load(f)\n    tmp = defaultdict(lambda: defaultdict(list))\n    n = 0\n    for root in roots:\n        for p in root.rglob(\"*\"):\n            if not p.is_file() or p.suffix.lower() not in {\".dcm\", \".dicom\"}:\n                continue\n            sid = next((part for part in p.parts if part in study_ids), None)\n            if sid is None:\n                continue\n            tmp[sid][str(p.parent)].append(str(p))\n            n += 1\n    out = {sid: [sorted(v) for _, v in groups.items()] for sid, groups in tmp.items()}\n    with open(INDEX_PATH, \"wb\") as f:\n        pickle.dump(out, f, protocol=pickle.HIGHEST_PROTOCOL)\n    print(\"indexed files:\", n, \"studies:\", len(out), \"series:\", sum(len(v) for v in out.values()))\n    return out\n\nseries_index = build_series_index(DICOM_ROOTS, ALL_IDS, rebuild=REBUILD_INDEX)\nif not series_index:\n    raise RuntimeError(\"No DICOM indexed under the fixed Kaggle competition directory.\")\n\ndef safe_attr(ds, name, default=None):\n    try:\n        value = getattr(ds, name)\n        return default if value is None else value\n    except Exception:\n        return default\n\ndef float_list(value):\n    try:\n        return [float(x) for x in value]\n    except Exception:\n        return None\n\ndef infer_plane(iop, desc):\n    if iop is not None and len(iop) >= 6:\n        try:\n            normal = np.cross(np.asarray(iop[:3]), np.asarray(iop[3:6]))\n            return {0: \"sagittal\", 1: \"coronal\", 2: \"axial\"}[int(np.argmax(np.abs(normal)))]\n        except Exception:\n            pass\n    text = str(desc).lower()\n    if \"sag\" in text:\n        return \"sagittal\"\n    if \"cor\" in text:\n        return \"coronal\"\n    if \"ax\" in text or \"trans\" in text:\n        return \"axial\"\n    return \"unknown\"\n\ndef read_header(files):\n    import pydicom\n    try:\n        ds = pydicom.dcmread(files[0], stop_before_pixels=True, force=True)\n        desc = \" \".join(str(safe_attr(ds, k, \"\")) for k in [\"SeriesDescription\", \"ProtocolName\", \"SequenceName\"])\n        text = desc.lower()\n        return {\n            \"files\": files, \"n_slices\": len(files), \"description\": desc,\n            \"plane\": infer_plane(float_list(safe_attr(ds, \"ImageOrientationPatient\", None)), desc),\n            \"fluid\": any(x in text for x in [\"t2\", \"pd\", \"stir\", \"fs\", \"fat sat\", \"fatsat\", \"fluid\"]),\n            \"fat_sat\": any(x in text for x in [\"fs\", \"fat sat\", \"fatsat\", \"stir\", \"spair\", \"water\"]),\n            \"localizer\": any(x in text for x in [\"localizer\", \"scout\", \"loc\"]),\n        }\n    except Exception:\n        return {\"files\": files, \"n_slices\": len(files), \"description\": \"\", \"plane\": \"unknown\", \"fluid\": False, \"fat_sat\": False, \"localizer\": False}\n\nMETA_PATH = WORK_DIR / \"series_meta.pkl\"\n\ndef build_meta(index, rebuild=False):\n    if META_PATH.exists() and not rebuild:\n        with open(META_PATH, \"rb\") as f:\n            return pickle.load(f)\n    out = {}\n    for sid, groups in tqdm(index.items(), desc=\"series headers\"):\n        out[sid] = [read_header(files) for files in groups]\n    with open(META_PATH, \"wb\") as f:\n        pickle.dump(out, f, protocol=pickle.HIGHEST_PROTOCOL)\n    return out\n\nseries_meta = build_meta(series_index, rebuild=REBUILD_INDEX)\nprint(\"metadata studies:\", len(series_meta))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:05:07.85786Z","iopub.execute_input":"2026-09-08T15:05:07.858501Z","iopub.status.idle":"2026-09-08T15:17:02.759192Z","shell.execute_reply.started":"2026-09-08T15:05:07.858477Z","shell.execute_reply":"2026-09-08T15:17:02.758209Z"}},"outputs":[{"name":"stdout","text":"indexed files: 819635 studies: 4410 series: 24386\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"series headers:   0%|          | 0/4410 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"7654d71548e9401cb3585b9da69a0150"}},"metadata":{}},{"name":"stdout","text":"metadata studies: 4410\n","output_type":"stream"}],"execution_count":4},{"id":"ed0724f4-421e-4ed1-b412-3a878559501c","cell_type":"code","source":"SLOT_NAMES = [\"sagittal\", \"coronal\", \"axial\"]\n\ndef series_quality(r):\n    n = r[\"n_slices\"]\n    return 6 * int(r[\"fluid\"]) + 3 * int(r[\"fat_sat\"]) + 2 * int(16 <= n <= 80) + min(n, 64) / 64 - 8 * int(r[\"localizer\"]) - 2 * int(n < 5)\n\ndef select_series(study_id):\n    records = series_meta.get(str(study_id), [])\n    if not records:\n        return [None] * N_SLOTS\n    for r in records:\n        r[\"quality\"] = series_quality(r)\n    chosen, used = [], set()\n    for plane in SLOT_NAMES:\n        candidates = sorted([r for r in records if r[\"plane\"] == plane], key=lambda r: -r[\"quality\"])\n        r = candidates[0] if candidates else None\n        chosen.append(r)\n        if r is not None:\n            used.add(id(r))\n    remaining = sorted([r for r in records if id(r) not in used], key=lambda r: -r[\"quality\"])\n    for i in range(N_SLOTS):\n        if chosen[i] is None and remaining:\n            chosen[i] = remaining.pop(0)\n    return chosen\n\nids = [x for x in train_df[\"StudyInstanceUID\"].head(100) if x in series_meta]\nif ids:\n    print(\"example study:\", ids[0])\n    for i, r in enumerate(select_series(ids[0])):\n        print(i, SLOT_NAMES[i], None if r is None else (r[\"plane\"], r[\"n_slices\"], r[\"description\"][:70]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:17:02.760383Z","iopub.execute_input":"2026-09-08T15:17:02.760765Z","iopub.status.idle":"2026-09-08T15:17:02.770373Z","shell.execute_reply.started":"2026-09-08T15:17:02.76074Z","shell.execute_reply":"2026-09-08T15:17:02.769564Z"}},"outputs":[{"name":"stdout","text":"example study: 1.2.826.0.1.3680043.8.498.10004873229099053869093324292195817260\n0 sagittal ('sagittal', 22, 'DP CS_SAG  ')\n1 coronal ('coronal', 16, 'DP SPAIR_CS  ')\n2 axial ('axial', 18, 'STIR CS_ AX  ')\n","output_type":"stream"}],"execution_count":5},{"id":"0b680bc4-74c4-4bdb-9bb4-eee44b4bfd7f","cell_type":"markdown","source":"## 4. DICOM preprocessing and cache\n\nEach study is represented as a uint8 tensor with shape 3 x 32 x 224 x 224. The Dataset samples K_SLICES slices from each selected series.","metadata":{}},{"id":"749ec2ff-593d-41b4-ab97-3f1e45e473f0","cell_type":"code","source":"from PIL import Image\n\ndef resize_gray(arr, size=IMG_SIZE):\n    arr = np.asarray(arr, dtype=np.uint8)\n    try:\n        import cv2\n        return cv2.resize(arr, (size, size), interpolation=cv2.INTER_AREA)\n    except Exception:\n        resampling = getattr(Image, \"Resampling\", Image).BILINEAR\n        return np.asarray(Image.fromarray(arr).resize((size, size), resampling), dtype=np.uint8)\n\ndef sort_key(ds, normal):\n    try:\n        pos = np.asarray(ds.ImagePositionPatient, dtype=np.float32)\n        if normal is not None:\n            return (0, float(np.dot(pos, normal)))\n    except Exception:\n        pass\n    try:\n        return (1, float(ds.InstanceNumber))\n    except Exception:\n        return (2, 0)\n\ndef read_series(files):\n    import pydicom\n    items, normal = [], None\n    for path in files:\n        try:\n            ds = pydicom.dcmread(path, force=True)\n            iop = float_list(safe_attr(ds, \"ImageOrientationPatient\", None))\n            if normal is None and iop is not None and len(iop) >= 6:\n                normal = np.cross(np.asarray(iop[:3], dtype=np.float32), np.asarray(iop[3:6], dtype=np.float32))\n            arr = ds.pixel_array.astype(np.float32)\n            arr = arr * float(safe_attr(ds, \"RescaleSlope\", 1.0) or 1.0) + float(safe_attr(ds, \"RescaleIntercept\", 0.0) or 0.0)\n            if str(safe_attr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n                arr = np.max(arr) - arr\n            items.append((sort_key(ds, normal), arr))\n        except Exception:\n            continue\n    if not items:\n        return np.zeros((MAX_SLICES, IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n    items.sort(key=lambda x: x[0])\n    arrays = [x[1] for x in items]\n    flat = np.concatenate([a.reshape(-1) for a in arrays])\n    lo, hi = np.percentile(flat, [1, 99])\n    if not np.isfinite(lo) or not np.isfinite(hi) or hi <= lo:\n        lo, hi = float(flat.min()), float(flat.max() + 1e-5)\n    out = []\n    for arr in arrays:\n        arr = np.clip((arr - lo) / (hi - lo + 1e-6), 0, 1)\n        out.append(resize_gray((arr * 255).astype(np.uint8)))\n    volume = np.stack(out)\n    idx = np.rint(np.linspace(0, len(volume) - 1, MAX_SLICES)).astype(int)\n    return volume[idx]\n\ndef preprocess_study(study_id):\n    x = np.zeros((N_SLOTS, MAX_SLICES, IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n    mask = np.zeros(N_SLOTS, dtype=np.bool_)\n    for slot, record in enumerate(select_series(study_id)):\n        if record is not None:\n            x[slot] = read_series(record[\"files\"])\n            mask[slot] = True\n    return x, mask\n\ndef cache_path(study_id):\n    return CACHE_DIR / (re.sub(r\"[^A-Za-z0-9_.-]\", \"_\", str(study_id)) + \".npz\")\n\ndef cache_one(study_id, force=False):\n    path = cache_path(study_id)\n    if path.exists() and not force:\n        return path\n    x, mask = preprocess_study(study_id)\n    np.savez_compressed(path, x=x, slot_mask=mask)\n    return path\n\ndef prepare_cache(ids):\n    ids = list(dict.fromkeys(map(str, ids)))\n    if MAX_CACHE_STUDIES is not None:\n        ids = ids[:MAX_CACHE_STUDIES]\n    for sid in tqdm(ids, desc=\"precompute cache\"):\n        cache_one(sid, force=REBUILD_CACHE)\n    print(\"cache files:\", len(list(CACHE_DIR.glob(\"*.npz\"))))\n\nif PREPARE_CACHE:\n    prepare_cache(list(train_df[\"StudyInstanceUID\"]) + list(test_df[\"StudyInstanceUID\"]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T15:17:02.771376Z","iopub.execute_input":"2026-09-08T15:17:02.77181Z","iopub.status.idle":"2026-09-08T17:46:20.590982Z","shell.execute_reply.started":"2026-09-08T15:17:02.771788Z","shell.execute_reply":"2026-09-08T17:46:20.590302Z"}},"outputs":[{"output_type":"display_data","data":{"text/plain":"precompute cache:   0%|          | 0/4410 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"d355b95e76714680b3c190a588de27b1"}},"metadata":{}},{"name":"stdout","text":"cache files: 4410\n","output_type":"stream"}],"execution_count":6},{"id":"e895a21e-0764-40de-8ddc-47cf03bd1415","cell_type":"markdown","source":"## 5. Dataset and augmentation\n\nOnly mild intensity augmentation is used. Horizontal flip is intentionally excluded because laterality matters for several targets.","metadata":{}},{"id":"ca5ac94e-3631-47d5-a4d4-593026e0122a","cell_type":"code","source":"def choose_indices(train_mode):\n    if K_SLICES >= MAX_SLICES:\n        return np.arange(MAX_SLICES)\n    idx = np.linspace(0, MAX_SLICES - 1, K_SLICES).round().astype(int)\n    if train_mode:\n        idx = np.clip(idx + np.random.choice([-1, 0, 1], size=K_SLICES, p=[.15, .7, .15]), 0, MAX_SLICES - 1)\n    return idx\n\ndef intensity_aug(x):\n    if np.random.rand() < .5:\n        x = np.clip(x * np.random.uniform(.9, 1.1) + np.random.uniform(-.04, .04), 0, 1)\n    if np.random.rand() < .15:\n        x = np.clip(x + np.random.normal(0, .012, x.shape).astype(\"float32\"), 0, 1)\n    return x\n\ndef load_study(study_id):\n    p = cache_path(study_id)\n    if p.exists():\n        z = np.load(p)\n        return z[\"x\"], z[\"slot_mask\"].astype(bool)\n    return preprocess_study(study_id)\n\nclass KneeDataset(Dataset):\n    def __init__(self, frame, train_mode=False):\n        self.frame = frame.reset_index(drop=True)\n        self.train_mode = train_mode\n        self.ids = self.frame[\"StudyInstanceUID\"].astype(str).tolist()\n        self.has_target = all(c in self.frame.columns for c in TARGET_COLUMNS)\n\n    def __len__(self):\n        return len(self.frame)\n\n    def __getitem__(self, i):\n        sid = self.ids[i]\n        x, slot_mask = load_study(sid)\n        x = x[:, choose_indices(self.train_mode)].astype(\"float32\") / 255.0\n        if self.train_mode:\n            x = intensity_aug(x)\n        item = {\"image\": torch.from_numpy(x), \"slot_mask\": torch.from_numpy(slot_mask), \"uid\": sid}\n        if self.has_target:\n            y = self.frame.loc[i, TARGET_COLUMNS].to_numpy(dtype=\"float32\")\n            item[\"target\"] = torch.from_numpy(y * (1 - LABEL_SMOOTHING) + .5 * LABEL_SMOOTHING)\n            w = self.frame.loc[i, TARGET_WEIGHT_COLUMNS].to_numpy(dtype=\"float32\")\n            item[\"weight\"] = torch.from_numpy(w)\n        return item\n\ndef make_loader(frame, train_mode, shuffle):\n    return DataLoader(\n        KneeDataset(frame, train_mode),\n        batch_size=BATCH_SIZE,\n        shuffle=shuffle,\n        num_workers=NUM_WORKERS,\n        pin_memory=True,\n        persistent_workers=NUM_WORKERS > 0,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:46:20.592177Z","iopub.execute_input":"2026-09-08T17:46:20.5935Z","iopub.status.idle":"2026-09-08T17:46:20.60981Z","shell.execute_reply.started":"2026-09-08T17:46:20.593472Z","shell.execute_reply":"2026-09-08T17:46:20.609051Z"}},"outputs":[],"execution_count":7},{"id":"941590a6-0d25-422a-8228-c79fa20a2792","cell_type":"markdown","source":"## 6. Slice encoder and study-level attention MIL","metadata":{}},{"id":"dac71b61-ddfd-4216-ae3d-d9ffd10f44ea","cell_type":"code","source":"from torchvision.models import resnet18, ResNet18_Weights\n\nclass ResNetBackbone(nn.Module):\n    def __init__(self, pretrained=True):\n        super().__init__()\n        try:\n            net = resnet18(weights=ResNet18_Weights.DEFAULT if pretrained else None)\n        except Exception as e:\n            print(\"ImageNet weights unavailable; random ResNet18:\", repr(e))\n            net = resnet18(weights=None)\n        self.body = nn.Sequential(*list(net.children())[:-1])\n        self.out_dim = 512\n\n    def forward(self, x):\n        return self.body(x).flatten(1)\n\nclass DINOv2Backbone(nn.Module):\n    def __init__(self, pretrained=True):\n        super().__init__()\n        import timm\n        try:\n            self.net = timm.create_model(\"vit_small_patch14_dinov2.lvd142m\", pretrained=pretrained, num_classes=0, img_size=IMG_SIZE, attn_implementation=ATTENTION_IMPLEMENTATION)\n        except TypeError:\n            self.net = timm.create_model(\"vit_small_patch14_dinov2.lvd142m\", pretrained=pretrained, num_classes=0, img_size=IMG_SIZE)\n        self.out_dim = int(getattr(self.net, \"num_features\", 384))\n\n    def forward(self, x):\n        z = self.net.forward_features(x)\n        if isinstance(z, dict):\n            z = z.get(\"x_norm_clstoken\", next(iter(z.values())))\n        return z[:, 0] if z.ndim == 3 else z.flatten(1)\n\ndef build_backbone():\n    if BACKBONE.lower() == \"dinov2\":\n        try:\n            return DINOv2Backbone(PRETRAINED)\n        except Exception as e:\n            print(\"DINOv2 unavailable; fallback to ResNet18:\", repr(e))\n    return ResNetBackbone(PRETRAINED)\n\nclass StudyMIL(nn.Module):\n    def __init__(self, n_targets):\n        super().__init__()\n        self.backbone = build_backbone()\n        d = self.backbone.out_dim\n        self.register_buffer(\"mean\", torch.tensor([.485, .456, .406]).view(1, 3, 1, 1))\n        self.register_buffer(\"std\", torch.tensor([.229, .224, .225]).view(1, 3, 1, 1))\n        self.attn = nn.Sequential(nn.Linear(d, 128), nn.Tanh(), nn.Dropout(.1), nn.Linear(128, 1))\n        self.slot_embedding = nn.Parameter(torch.randn(N_SLOTS, d) * .02)\n        self.head = nn.Sequential(nn.LayerNorm(N_SLOTS * d), nn.Linear(N_SLOTS * d, 512), nn.GELU(), nn.Dropout(.25), nn.Linear(512, n_targets))\n\n    def forward(self, image, slot_mask):\n        b, s, k, h, w = image.shape\n        x = image.reshape(b * s * k, 1, h, w).repeat(1, 3, 1, 1)\n        x = (x - self.mean.to(x.dtype)) / self.std.to(x.dtype)\n        z = self.backbone(x).reshape(b, s, k, -1)\n        z = z + self.slot_embedding[None, :, None, :].to(z.dtype)\n        scores = self.attn(z).squeeze(-1)\n        valid = slot_mask.bool().unsqueeze(-1).expand(-1, -1, k)\n        scores = scores.masked_fill(~valid, -1e4)\n        weights = torch.softmax(scores, dim=-1)\n        pooled = (weights.unsqueeze(-1) * z).sum(2)\n        pooled = pooled * slot_mask.to(pooled.dtype).unsqueeze(-1)\n        return self.head(pooled.reshape(b, -1))\n\ndef unwrap(model):\n    return model.module if isinstance(model, nn.DataParallel) else model\n\ndef make_model():\n    model = StudyMIL(len(TARGET_COLUMNS)).to(DEVICE)\n    if USE_DATA_PARALLEL and torch.cuda.device_count() > 1:\n        model = nn.DataParallel(model)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:46:20.610769Z","iopub.execute_input":"2026-09-08T17:46:20.610984Z","iopub.status.idle":"2026-09-08T17:46:27.319201Z","shell.execute_reply.started":"2026-09-08T17:46:20.610964Z","shell.execute_reply":"2026-09-08T17:46:27.318485Z"}},"outputs":[],"execution_count":8},{"id":"739881ea-b9cc-48d1-94b9-57f365178945","cell_type":"markdown","source":"## 7. Loss, macro ROC-AUC and training loop","metadata":{}},{"id":"42be3814-0490-4cc2-bc88-df6cdc4ec006","cell_type":"code","source":"AMP_ENABLED = USE_AMP and DEVICE == \"cuda\"\n\ndef move_batch(batch):\n    image = batch[\"image\"].to(DEVICE, non_blocking=True)\n    mask = batch[\"slot_mask\"].to(DEVICE, non_blocking=True)\n    target = batch.get(\"target\")\n    weight = batch.get(\"weight\")\n    if target is not None:\n        target = target.to(DEVICE, non_blocking=True)\n    if weight is not None:\n        weight = weight.to(DEVICE, non_blocking=True)\n    return image, mask, target, weight\n\ndef macro_auc(y, p):\n    scores, detail = [], {}\n    for j, name in enumerate(TARGET_COLUMNS):\n        if len(np.unique(y[:, j])) < 2:\n            detail[name] = np.nan\n        else:\n            detail[name] = float(roc_auc_score(y[:, j], p[:, j]))\n            scores.append(detail[name])\n    return (float(np.mean(scores)) if scores else .5), detail\n\ndef train_epoch(model, loader, optimizer, scaler):\n    model.train()\n    losses = []\n    optimizer.zero_grad(set_to_none=True)\n    for step, batch in enumerate(tqdm(loader, desc=\"train\", leave=False)):\n        image, mask, target, weight = move_batch(batch)\n        with torch.cuda.amp.autocast(enabled=AMP_ENABLED, dtype=torch.float16):\n            logits = model(image, mask)\n            element_loss = F.binary_cross_entropy_with_logits(logits, target, reduction=\"none\")\n            loss = (element_loss * weight).sum() / weight.sum().clamp_min(1.0)\n        scaler.scale(loss / ACCUM_STEPS).backward()\n        if (step + 1) % ACCUM_STEPS == 0 or step + 1 == len(loader):\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n        losses.append(float(loss.detach().cpu()))\n    return float(np.mean(losses))\n\n@torch.no_grad()\ndef predict_loader(model, loader):\n    model.eval()\n    preds, trues, ids = [], [], []\n    for batch in tqdm(loader, desc=\"valid\", leave=False):\n        image, mask, target, weight = move_batch(batch)\n        with torch.cuda.amp.autocast(enabled=AMP_ENABLED, dtype=torch.float16):\n            p = torch.sigmoid(model(image, mask)).float().cpu().numpy()\n        preds.append(p)\n        ids.extend(batch[\"uid\"])\n        if target is not None:\n            trues.append(target.float().cpu().numpy())\n    return ids, np.concatenate(preds), np.concatenate(trues) if trues else None\n\ndef save_ckpt(model, optimizer, scheduler, fold, epoch, score, path):\n    torch.save({\"model\": unwrap(model).state_dict(), \"optimizer\": optimizer.state_dict(), \"scheduler\": scheduler.state_dict(), \"fold\": fold, \"epoch\": epoch, \"score\": score}, path)\n\ndef load_ckpt(model, path):\n    state = torch.load(path, map_location=\"cpu\")\n    unwrap(model).load_state_dict(state[\"model\"], strict=True)\n    return state","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:46:27.32Z","iopub.execute_input":"2026-09-08T17:46:27.320398Z","iopub.status.idle":"2026-09-08T17:46:27.335368Z","shell.execute_reply.started":"2026-09-08T17:46:27.320375Z","shell.execute_reply":"2026-09-08T17:46:27.334373Z"}},"outputs":[],"execution_count":9},{"id":"4cba38f9-f527-4c58-a060-5541c4fc6b3d","cell_type":"markdown","source":"## 8. Smoke test","metadata":{}},{"id":"d1afcc60-8efc-43e2-a87b-924641991592","cell_type":"code","source":"if RUN_SMOKE_TEST:\n    frame = train_df.head(min(2, len(train_df)))\n    loader = make_loader(frame, True, False)\n    model = make_model()\n    batch = next(iter(loader))\n    image, mask, target, weight = move_batch(batch)\n    with torch.cuda.amp.autocast(enabled=AMP_ENABLED, dtype=torch.float16):\n        logits = model(image, mask)\n        element_loss = F.binary_cross_entropy_with_logits(logits, target, reduction=\"none\")\n        loss = (element_loss * weight).sum() / weight.sum().clamp_min(1.0)\n    loss.backward()\n    assert tuple(logits.shape) == (len(frame), len(TARGET_COLUMNS))\n    print(\"SMOKE TEST PASSED\", tuple(image.shape), tuple(logits.shape), float(loss.detach().cpu()))\n    del model, loader, batch, image, mask, target, weight, logits\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:46:27.336482Z","iopub.execute_input":"2026-09-08T17:46:27.33698Z","iopub.status.idle":"2026-09-08T17:48:22.690464Z","shell.execute_reply.started":"2026-09-08T17:46:27.33694Z","shell.execute_reply":"2026-09-08T17:48:22.689512Z"}},"outputs":[{"name":"stderr","text":"'[Errno -3] Temporary failure in name resolution' thrown while requesting HEAD https://huggingface.co/timm/vit_small_patch14_dinov2.lvd142m/resolve/main/model.safetensors\nRetrying in 1s [Retry 1/5].\n","output_type":"stream"},{"name":"stdout","text":"DINOv2 unavailable; fallback to ResNet18: RuntimeError('Cannot send a request, as the client has been closed.')\nDownloading: \"https://download.pytorch.org/models/resnet18-f37072fd.pth\" to /root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth\nImageNet weights unavailable; random ResNet18: URLError(gaierror(-3, 'Temporary failure in name resolution'))\nSMOKE TEST PASSED (2, 3, 12, 224, 224) (2, 12) 0.27569779753685\n","output_type":"stream"}],"execution_count":10},{"id":"14ca2cea-f0f2-443d-8fe6-44aadded7b22","cell_type":"markdown","source":"## 9. Five-fold study-level training","metadata":{}},{"id":"05bfa437-8f14-4087-92eb-114f66b22062","cell_type":"code","source":"gkf = GroupKFold(n_splits=N_FOLDS)\nsplits = list(gkf.split(train_df, groups=train_df[\"StudyInstanceUID\"]))\nOOF = np.zeros((len(train_df), len(TARGET_COLUMNS)), dtype=\"float32\")\nfold_rows = []\n\ndef run_fold(fold, train_idx, valid_idx):\n    seed_everything(SEED + fold)\n    tr = train_df.iloc[train_idx].reset_index(drop=True)\n    va = train_df.iloc[valid_idx].reset_index(drop=True)\n    tr_loader = make_loader(tr, True, True)\n    va_loader = make_loader(va, False, False)\n    model = make_model()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = torch.cuda.amp.GradScaler(enabled=AMP_ENABLED)\n    path = CKPT_DIR / f\"fold_{fold}.pt\"\n\n    if not (RESUME and path.exists()):\n        best, stale = -np.inf, 0\n        for epoch in range(EPOCHS):\n            t0 = time.time()\n            loss = train_epoch(model, tr_loader, optimizer, scaler)\n            _, pred, true = predict_loader(model, va_loader)\n            score, detail = macro_auc(true, pred)\n            scheduler.step()\n            print(f\"fold={fold} epoch={epoch+1}/{EPOCHS} loss={loss:.5f} auc={score:.5f} time={time.time()-t0:.1f}s\")\n            if score > best:\n                best, stale = score, 0\n                save_ckpt(model, optimizer, scheduler, fold, epoch, score, path)\n            else:\n                stale += 1\n            if stale >= PATIENCE:\n                print(\"early stopping\")\n                break\n    else:\n        print(\"resume:\", path)\n\n    load_ckpt(model, path)\n    _, pred, true = predict_loader(model, va_loader)\n    score, detail = macro_auc(true, pred)\n    print(\"best fold score:\", fold, score)\n    del tr_loader, va_loader, model, optimizer, scheduler, scaler\n    torch.cuda.empty_cache()\n    gc.collect()\n    return valid_idx, pred, score, detail\n\nif RUN_FULL_TRAIN:\n    for fold, (tr_idx, va_idx) in enumerate(splits):\n        if fold not in FOLDS_TO_RUN:\n            continue\n        idx, pred, score, detail = run_fold(fold, tr_idx, va_idx)\n        OOF[idx] = pred\n        fold_rows.append({\"fold\": fold, \"macro_auc\": score, **detail})\n        pd.DataFrame(fold_rows).to_csv(WORK_DIR / \"fold_scores.csv\", index=False)\n\n    oof_score, oof_detail = macro_auc(train_df[TARGET_COLUMNS].to_numpy(), OOF)\n    oof_df = pd.DataFrame({\"StudyInstanceUID\": train_df[\"StudyInstanceUID\"]})\n    for j, c in enumerate(TARGET_COLUMNS):\n        oof_df[c] = OOF[:, j]\n    oof_df.to_csv(WORK_DIR / \"oof.csv\", index=False)\n    print(\"OOF macro AUC:\", oof_score)\n    display(pd.DataFrame({\"target\": list(oof_detail), \"auc\": list(oof_detail.values())}))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:48:22.69188Z","iopub.execute_input":"2026-09-08T17:48:22.692854Z","iopub.status.idle":"2026-09-08T17:59:09.21358Z","shell.execute_reply.started":"2026-09-08T17:48:22.692823Z","shell.execute_reply":"2026-09-08T17:59:09.21218Z"}},"outputs":[{"name":"stderr","text":"'[Errno -3] Temporary failure in name resolution' thrown while requesting HEAD https://huggingface.co/timm/vit_small_patch14_dinov2.lvd142m/resolve/main/model.safetensors\nRetrying in 1s [Retry 1/5].\n","output_type":"stream"},{"name":"stdout","text":"DINOv2 unavailable; fallback to ResNet18: RuntimeError('Cannot send a request, as the client has been closed.')\nDownloading: \"https://download.pytorch.org/models/resnet18-f37072fd.pth\" to /root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth\nImageNet weights unavailable; random ResNet18: URLError(gaierror(-3, 'Temporary failure in name resolution'))\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"train:   0%|          | 0/882 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":""}},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"valid:   0%|          | 0/221 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":""}},"metadata":{}},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mValueError\u001b[0m                                Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_84/1548929147.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m     49\u001b[0m         \u001b[0;32mif\u001b[0m \u001b[0mfold\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mFOLDS_TO_RUN\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     50\u001b[0m             \u001b[0;32mcontinue\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 51\u001b[0;31m         \u001b[0midx\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mpred\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mscore\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdetail\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrun_fold\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mfold\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtr_idx\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mva_idx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     52\u001b[0m         \u001b[0mOOF\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0midx\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpred\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     53\u001b[0m         \u001b[0mfold_rows\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mappend\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m{\u001b[0m\u001b[0;34m\"fold\"\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mfold\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"macro_auc\"\u001b[0m\u001b[0;34m:\u001b[0m \u001b[0mscore\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mdetail\u001b[0m\u001b[0;34m}\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_84/1548929147.py\u001b[0m in \u001b[0;36mrun_fold\u001b[0;34m(fold, train_idx, valid_idx)\u001b[0m\n\u001b[1;32m     22\u001b[0m             \u001b[0mloss\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtrain_epoch\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtr_loader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0moptimizer\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mscaler\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     23\u001b[0m             \u001b[0m_\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mpred\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtrue\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpredict_loader\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mva_loader\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 24\u001b[0;31m             \u001b[0mscore\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdetail\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mmacro_auc\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mtrue\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mpred\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     25\u001b[0m             \u001b[0mscheduler\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstep\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     26\u001b[0m             \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34mf\"fold={fold} epoch={epoch+1}/{EPOCHS} loss={loss:.5f} auc={score:.5f} time={time.time()-t0:.1f}s\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_84/2653445340.py\u001b[0m in \u001b[0;36mmacro_auc\u001b[0;34m(y, p)\u001b[0m\n\u001b[1;32m     18\u001b[0m             \u001b[0mdetail\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mname\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mnp\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mnan\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     19\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 20\u001b[0;31m             \u001b[0mdetail\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mname\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfloat\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mroc_auc_score\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mj\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mp\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mj\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     21\u001b[0m             \u001b[0mscores\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mappend\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdetail\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mname\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     22\u001b[0m     \u001b[0;32mreturn\u001b[0m \u001b[0;34m(\u001b[0m\u001b[0mfloat\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mnp\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmean\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mscores\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;32mif\u001b[0m \u001b[0mscores\u001b[0m \u001b[0;32melse\u001b[0m \u001b[0;36m.5\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdetail\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/sklearn/utils/_param_validation.py\u001b[0m in \u001b[0;36mwrapper\u001b[0;34m(*args, **kwargs)\u001b[0m\n\u001b[1;32m    214\u001b[0m                     )\n\u001b[1;32m    215\u001b[0m                 ):\n\u001b[0;32m--> 216\u001b[0;31m                     \u001b[0;32mreturn\u001b[0m \u001b[0mfunc\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    217\u001b[0m             \u001b[0;32mexcept\u001b[0m \u001b[0mInvalidParameterError\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0me\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    218\u001b[0m                 \u001b[0;31m# When the function is just a wrapper around an estimator, we allow\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/sklearn/metrics/_ranking.py\u001b[0m in \u001b[0;36mroc_auc_score\u001b[0;34m(y_true, y_score, average, sample_weight, max_fpr, multi_class, labels)\u001b[0m\n\u001b[1;32m    647\u001b[0m         )\n\u001b[1;32m    648\u001b[0m     \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m  \u001b[0;31m# multilabel-indicator\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 649\u001b[0;31m         return _average_binary_score(\n\u001b[0m\u001b[1;32m    650\u001b[0m             \u001b[0mpartial\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0m_binary_roc_auc_score\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mmax_fpr\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mmax_fpr\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    651\u001b[0m             \u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.12/dist-packages/sklearn/metrics/_base.py\u001b[0m in \u001b[0;36m_average_binary_score\u001b[0;34m(binary_metric, y_true, y_score, average, sample_weight)\u001b[0m\n\u001b[1;32m     64\u001b[0m     \u001b[0my_type\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtype_of_target\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     65\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0my_type\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0;32min\u001b[0m \u001b[0;34m(\u001b[0m\u001b[0;34m\"binary\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m\"multilabel-indicator\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 66\u001b[0;31m         \u001b[0;32mraise\u001b[0m \u001b[0mValueError\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"{0} format is not supported\"\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mformat\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_type\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     67\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     68\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0my_type\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0;34m\"binary\"\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mValueError\u001b[0m: continuous format is not supported"],"ename":"ValueError","evalue":"continuous format is not supported","output_type":"error"}],"execution_count":11},{"id":"8a4276e9-f42d-4265-a8f9-bde7443412ec","cell_type":"markdown","source":"## 10. Test inference and submission","metadata":{}},{"id":"5f22d7eb-f92c-4b43-9bdc-babb7251c264","cell_type":"code","source":"@torch.no_grad()\ndef infer_checkpoint(frame, checkpoint_path):\n    loader = make_loader(frame, False, False)\n    model = make_model()\n    load_ckpt(model, checkpoint_path)\n    model.eval()\n    preds, ids = [], []\n    for batch in tqdm(loader, desc=f\"infer {checkpoint_path.name}\"):\n        image, mask, _, _ = move_batch(batch)\n        with torch.cuda.amp.autocast(enabled=AMP_ENABLED, dtype=torch.float16):\n            p = torch.sigmoid(model(image, mask))\n            if USE_TTA:\n                p = (p + torch.sigmoid(model(torch.clamp(image * .96 + .02, 0, 1), mask))) / 2\n        preds.append(p.float().cpu().numpy())\n        ids.extend(batch[\"uid\"])\n    del loader, model\n    torch.cuda.empty_cache()\n    gc.collect()\n    return ids, np.concatenate(preds)\n\nif RUN_TEST_INFERENCE:\n    paths = [CKPT_DIR / f\"fold_{f}.pt\" for f in FOLDS_TO_RUN if (CKPT_DIR / f\"fold_{f}.pt\").exists()]\n    if not paths:\n        raise FileNotFoundError(\"No fold checkpoint found. Run training first.\")\n    all_pred, test_ids = [], None\n    for path in paths:\n        ids, pred = infer_checkpoint(test_df, path)\n        test_ids = ids if test_ids is None else test_ids\n        all_pred.append(pred)\n    test_pred = np.mean(all_pred, axis=0).clip(1e-5, 1 - 1e-5)\n    pred_map = {str(uid): test_pred[i] for i, uid in enumerate(test_ids)}\n\n    if sample_raw is not None:\n        submission = sample_raw.copy()\n        sample_ids = norm_uid(submission[sample_id]).tolist()\n        for j, c in enumerate(TARGET_COLUMNS):\n            if c not in submission.columns:\n                submission[c] = .5\n            submission[c] = [pred_map.get(str(uid), np.repeat(.5, len(TARGET_COLUMNS)))[j] for uid in sample_ids]\n    else:\n        submission = pd.DataFrame({\"StudyInstanceUID\": test_ids})\n        for j, c in enumerate(TARGET_COLUMNS):\n            submission[c] = test_pred[:, j]\n\n    SUBMISSION_PATH = Path(\"/kaggle/working/submission.csv\")\n    submission.to_csv(SUBMISSION_PATH, index=False)\n    print(\"saved:\", SUBMISSION_PATH, \"shape:\", submission.shape)\n    display(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:59:09.214367Z","iopub.status.idle":"2026-09-08T17:59:09.214612Z","shell.execute_reply.started":"2026-09-08T17:59:09.214496Z","shell.execute_reply":"2026-09-08T17:59:09.214509Z"}},"outputs":[],"execution_count":null},{"id":"4ad15360-2d85-409f-a678-b4dd1ecc0e74","cell_type":"markdown","source":"## 11. Final submission check\n\nThe file to submit is /kaggle/working/submission.csv.","metadata":{}},{"id":"1b6094bc-76f2-45bc-9505-941a8e7ee9d1","cell_type":"code","source":"sub_path = Path(\"/kaggle/working/submission.csv\")\nif sub_path.exists():\n    sub = pd.read_csv(sub_path)\n    expected = len(sample_raw) if sample_raw is not None else len(test_df)\n    assert len(sub) == expected\n    assert all(c in sub.columns for c in TARGET_COLUMNS)\n    values = sub[TARGET_COLUMNS].to_numpy(dtype=float)\n    assert np.isfinite(values).all()\n    assert (values >= 0).all() and (values <= 1).all()\n    print(\"FINAL CHECK PASSED\", sub.shape)\n    display(sub[TARGET_COLUMNS].describe().T[[\"min\", \"max\", \"mean\"]])\nelse:\n    print(\"submission.csv not found.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-08T17:59:09.216749Z","iopub.status.idle":"2026-09-08T17:59:09.217114Z","shell.execute_reply.started":"2026-09-08T17:59:09.216932Z","shell.execute_reply":"2026-09-08T17:59:09.216953Z"}},"outputs":[],"execution_count":null},{"id":"bdfa6137-8ab0-40b9-837a-f8f1127c3664","cell_type":"markdown","source":"## Suggested improvements after the baseline works\n\n1. Change BACKBONE to dinov2 after attaching compatible timm weights.\n2. Increase K_SLICES to 16 or 24 if VRAM and runtime allow.\n3. Try target-specific attention or max pooling.\n4. Add report-derived soft labels only after validating their quality.\n5. Keep study-level splitting and avoid horizontal flips without laterality normalization.\n6. Checkpoints are saved under /kaggle/working/rsna_knee_abnormality_v1/checkpoints.","metadata":{}}]}