{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from pathlib import Path\nimport random, warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom PIL import Image\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import resnet18\nimport os, sys, time, json, math, random, argparse\nimport numpy as np, pandas as pd\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nimport timm\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\ndevice = 'cpu'\nlabels = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA', 'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', \"Baker's\", 'Contusion', 'Fracture']\nroot = next(p for p in Path('/kaggle/input').glob('**/train.csv') if (p.parent / 'sample_submission.csv').exists())\ntrain = pd.read_csv(root); test = pd.read_csv(root.parent / 'test.csv')\ntrain_series = pd.read_csv(root.parent / 'train_series.csv'); test_series = pd.read_csv(root.parent / 'test_series.csv')\ngold = train.dropna(subset=labels).reset_index(drop=True)\nprint(device, 'gold studies:', len(gold), 'test studies:', len(test))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:30.742733Z","iopub.execute_input":"2026-08-26T12:26:30.743006Z","iopub.status.idle":"2026-08-26T12:26:45.378909Z","shell.execute_reply.started":"2026-08-26T12:26:30.742974Z","shell.execute_reply":"2026-08-26T12:26:45.378181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pd.reset_option('display.max_colwidth')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.380113Z","iopub.execute_input":"2026-08-26T12:26:45.380332Z","iopub.status.idle":"2026-08-26T12:26:45.383782Z","shell.execute_reply.started":"2026-08-26T12:26:45.38031Z","shell.execute_reply":"2026-08-26T12:26:45.382956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# gold.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.384805Z","iopub.execute_input":"2026-08-26T12:26:45.38507Z","iopub.status.idle":"2026-08-26T12:26:45.577308Z","shell.execute_reply.started":"2026-08-26T12:26:45.38504Z","shell.execute_reply":"2026-08-26T12:26:45.576354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# gold.loc[[0,1,2,3], \"Report\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.578911Z","iopub.execute_input":"2026-08-26T12:26:45.579241Z","iopub.status.idle":"2026-08-26T12:26:45.608499Z","shell.execute_reply.started":"2026-08-26T12:26:45.579218Z","shell.execute_reply":"2026-08-26T12:26:45.607869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_series.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.609372Z","iopub.execute_input":"2026-08-26T12:26:45.609678Z","iopub.status.idle":"2026-08-26T12:26:45.644003Z","shell.execute_reply.started":"2026-08-26T12:26:45.609654Z","shell.execute_reply":"2026-08-26T12:26:45.643469Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !find /kaggle/input/datasets -maxdepth 3 -type f | head -100","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.644803Z","iopub.execute_input":"2026-08-26T12:26:45.645088Z","iopub.status.idle":"2026-08-26T12:26:45.6621Z","shell.execute_reply.started":"2026-08-26T12:26:45.645055Z","shell.execute_reply":"2026-08-26T12:26:45.661333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\npath = \"/kaggle/input/datasets/laymond/rsna-knee-abnormality-qwen3-8b-weak-labels/qwen_knee_weak_labels.csv\"\n\ndf = pd.read_csv(path)\n\nprint(\"Shape:\", df.shape)\nprint(\"\\nColumns:\")\nprint(df.columns.tolist())\n\ndisplay(df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:45.663048Z","iopub.execute_input":"2026-08-26T12:26:45.663306Z","iopub.status.idle":"2026-08-26T12:26:46.229442Z","shell.execute_reply.started":"2026-08-26T12:26:45.663277Z","shell.execute_reply":"2026-08-26T12:26:46.228702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# label_cols = [\n#     \"ACL__label\",\n#     \"MCL__label\",\n#     \"Medial Meniscus__label\",\n#     \"Lateral Meniscus__label\",\n#     \"Medial OA__label\",\n#     \"Lateral OA__label\",\n#     \"PF OA__label\",\n#     \"Effusion__label\",\n#     \"Synovitis__label\",\n#     \"Baker's__label\",\n#     \"Contusion__label\",\n#     \"Fracture__label\",\n# ]\n\n# for col in label_cols:\n#     print(f\"\\n{col}\")\n#     print(df[col].value_counts(dropna=False).head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:46.230391Z","iopub.execute_input":"2026-08-26T12:26:46.230791Z","iopub.status.idle":"2026-08-26T12:26:46.234945Z","shell.execute_reply.started":"2026-08-26T12:26:46.230767Z","shell.execute_reply":"2026-08-26T12:26:46.234176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# confidence_cols = [c for c in df.columns if c.endswith(\"__confidence\")]\n\n# display(df[confidence_cols].describe().T)\n\n# for col in confidence_cols:\n#     print(f\"\\n{col}\")\n#     print(df[col].describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:46.235895Z","iopub.execute_input":"2026-08-26T12:26:46.236142Z","iopub.status.idle":"2026-08-26T12:26:46.253916Z","shell.execute_reply.started":"2026-08-26T12:26:46.236119Z","shell.execute_reply":"2026-08-26T12:26:46.253135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\npath = \"/kaggle/input/datasets/laymond/rsna-knee-abnormality-qwen3-8b-weak-labels/qwen_knee_weak_labels.csv\"\n\ndf = pd.read_csv(path)\n\nfindings = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\ndef make_soft_target(row, finding):\n    label = row[f\"{finding}__label\"]\n    confidence = row[f\"{finding}__confidence\"]\n    negated = row[f\"{finding}__negated\"]\n\n    if pd.isna(label) or pd.isna(confidence):\n        return np.nan\n\n    label = float(label)\n    confidence = float(confidence)\n    negated = float(negated)\n\n    # Explicit negation → strong negative\n    if negated == 1 and label == 0:\n        return 0.05\n\n    # Positive prediction\n    if label == 1:\n        return confidence\n\n    # Negative prediction\n    return 1.0 - confidence\n\n\nfor finding in findings:\n    df[f\"{finding}__soft\"] = df.apply(\n        lambda row: make_soft_target(row, finding),\n        axis=1\n    )\n\nsoft_cols = [f\"{finding}__soft\" for finding in findings]\n\nprint(\"Before dropping:\", len(df))\n\ndf = df.dropna(subset=soft_cols).reset_index(drop=True)\n\nprint(\"After dropping:\", len(df))\nprint(\"Dropped:\", 4407 - len(df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:46.256526Z","iopub.execute_input":"2026-08-26T12:26:46.257137Z","iopub.status.idle":"2026-08-26T12:26:46.97632Z","shell.execute_reply.started":"2026-08-26T12:26:46.257113Z","shell.execute_reply":"2026-08-26T12:26:46.975655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display(\n    df[\n        [\n            \"ACL__label\",\n            \"ACL__confidence\",\n            \"ACL__mentioned\",\n            \"ACL__negated\",\n            \"ACL__soft\"\n        ]\n    ].head(20)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:46.977371Z","iopub.execute_input":"2026-08-26T12:26:46.977788Z","iopub.status.idle":"2026-08-26T12:26:46.991831Z","shell.execute_reply.started":"2026-08-26T12:26:46.977749Z","shell.execute_reply":"2026-08-26T12:26:46.991218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_data = df[[\"StudyInstanceUID\"]].copy()\n\nfor finding in findings:\n    label_data[finding] = df[f\"{finding}__soft\"]\n\ndisplay(label_data.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:46.992728Z","iopub.execute_input":"2026-08-26T12:26:46.993217Z","iopub.status.idle":"2026-08-26T12:26:47.022365Z","shell.execute_reply.started":"2026-08-26T12:26:46.993193Z","shell.execute_reply":"2026-08-26T12:26:47.021699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels_path = \"/kaggle/working/labels_soft.parquet\"\n\nlabel_data.to_parquet(\n    labels_path,\n    index=False\n)\n\nprint(f\"Saved to: {labels_path}\")\nprint(label_data.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.023442Z","iopub.execute_input":"2026-08-26T12:26:47.023752Z","iopub.status.idle":"2026-08-26T12:26:47.119089Z","shell.execute_reply.started":"2026-08-26T12:26:47.023723Z","shell.execute_reply":"2026-08-26T12:26:47.118321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bad = df[df[soft_cols].isna().any(axis=1)]\n\nprint(bad[\"StudyInstanceUID\"].tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.120021Z","iopub.execute_input":"2026-08-26T12:26:47.120215Z","iopub.status.idle":"2026-08-26T12:26:47.126765Z","shell.execute_reply.started":"2026-08-26T12:26:47.120196Z","shell.execute_reply":"2026-08-26T12:26:47.125852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, sys, time, json, math, random, argparse\nimport numpy as np, pandas as pd\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nimport timm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.127762Z","iopub.execute_input":"2026-08-26T12:26:47.128302Z","iopub.status.idle":"2026-08-26T12:26:47.137242Z","shell.execute_reply.started":"2026-08-26T12:26:47.128266Z","shell.execute_reply":"2026-08-26T12:26:47.136335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_data.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.138201Z","iopub.execute_input":"2026-08-26T12:26:47.138643Z","iopub.status.idle":"2026-08-26T12:26:47.157569Z","shell.execute_reply.started":"2026-08-26T12:26:47.138621Z","shell.execute_reply":"2026-08-26T12:26:47.156746Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\ndf = pd.read_parquet(\"/kaggle/working/labels_soft.parquet\")\n\nLAB = [\"ACL\",\"MCL\",\"Medial Meniscus\",\"Lateral Meniscus\",\"Medial OA\",\"Lateral OA\",\"PF OA\",\n       \"Effusion\",\"Synovitis\",\"Baker's\",\"Contusion\",\"Fracture\"]\n\nprint(df[LAB].isna().sum())\nprint(df[LAB].min())\nprint(df[LAB].max())\nprint(df[LAB].describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.158356Z","iopub.execute_input":"2026-08-26T12:26:47.158692Z","iopub.status.idle":"2026-08-26T12:26:47.25831Z","shell.execute_reply.started":"2026-08-26T12:26:47.158671Z","shell.execute_reply":"2026-08-26T12:26:47.257584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_parquet(\"/kaggle/working/labels_soft.parquet\")\n\nbad = df[df[LAB].isna().any(axis=1)]\n\nprint(bad[[\"StudyInstanceUID\"] + LAB].to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.259228Z","iopub.execute_input":"2026-08-26T12:26:47.25952Z","iopub.status.idle":"2026-08-26T12:26:47.271784Z","shell.execute_reply.started":"2026-08-26T12:26:47.259498Z","shell.execute_reply":"2026-08-26T12:26:47.271133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Bad IDs:\")\nprint(bad.index.tolist() if df.index.name == \"StudyInstanceUID\" else bad[\"StudyInstanceUID\"].tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.272829Z","iopub.execute_input":"2026-08-26T12:26:47.273268Z","iopub.status.idle":"2026-08-26T12:26:47.278341Z","shell.execute_reply.started":"2026-08-26T12:26:47.273243Z","shell.execute_reply":"2026-08-26T12:26:47.277574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"Knee MRI: training the twelve-finding model\n\nThis is the training half of the model behind the public 0.924 inference notebook. It reads the\nprecomputed slice stacks, trains a CoAtNet backbone with a per-finding attention pooling head, and\nwrites a checkpoint you can drop straight into that notebook.\n\nWHAT YOU NEED ATTACHED\n\n  1. The preprocessed corpus, both parts:\n       kaggle.com/datasets/dreaddevelopment/knee-raptor-corpus          (3,200 studies)\n       kaggle.com/datasets/dreaddevelopment/knee-raptor-corpus-ext      (1,207 studies)\n     Every study is already reduced to a fixed 44 x 336 x 336 uint8 stack, so no DICOM reading\n     happens here. The two parts concatenate in order.\n\n  2. The competition data, for train.csv.\n\n  3. Training labels, as a parquet with a StudyInstanceUID column and the twelve finding columns.\n     THIS IS NOT PROVIDED, and it is the one thing you have to bring yourself. See below.\n\nTHE LABEL PROBLEM, WHICH IS THE REAL PROBLEM\n\nThe competition gives you 4,407 studies and structured labels for only 58 of them. Every other\nstudy carries a free-text radiology report and nothing else. So before any of this trains, you need\nto turn 4,349 reports into twelve numbers each.\n\nThe approach behind the published weights was to read each report with a language model and emit\ntwelve probabilities rather than twelve yes or no answers: a report that hedges, saying a tear is\nsuspected, becomes something near 0.8 rather than a 1. Soft targets are far more forgiving than\nforcing every hedged sentence into a hard label, and the loss here expects them. The 58 studies\nthat come with real labels are held out and used only for validation, never trained on.\n\nPoint --labels at your own parquet built that way. The format is one row per study: a\nStudyInstanceUID column plus the twelve finding columns, values between 0 and 1.\n\nWHAT THE MODEL DOES\n\nThree neighbouring slices are stacked into the three channels of one image, so the network sees a\nlittle of what lies above and below the middle slice: most of the benefit of a 3D model at the cost\nof a 2D one. Each of these three-slice windows goes through the backbone, and the windows are then\npooled by an attention layer that has separate weights for each of the twelve findings. That last\npart matters more than anything else here. A cruciate tear may be visible on two slices while\nosteoarthritis spreads across many, and one shared pooling weight forces those to compete; giving\neach finding its own attention lets each draw on the slices that actually show it.\n\nTraining samples k windows per study at random and evaluates on k_eval windows spread evenly, so\neach epoch sees a different view of the same study. Nothing else is augmented.\n\nAt the end it keeps the best epoch by validation macro-AUC, and also writes a checkpoint that\naverages the weights of the best three epochs. Weight averaging costs nothing at inference, unlike\naveraging predictions from three models, and it usually gives a small gain.\n\nTYPICAL RUN\n\n  python train_knee.py --arch coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k --res 384 --epochs 16\n      --bs 8 --k 12 --k_eval 24 --grad_ckpt --tag mymodel --labels /kaggle/input/YOURS/labels.parquet\n\nAbout three hours on one 4090 for 16 epochs at 384. --grad_ckpt trades a little speed for a lot of\nmemory and is what makes bs 8 fit on a 24 GB card. Use --smoke for a fast wiring check.\n\nThe checkpoint it writes is a dict with keys model, arch, res and lab, which is exactly what the\ninference notebook expects.\n\nRESUMING ACROSS SESSIONS\n\n  Kaggle sessions have a hard time cap, so a long run can get killed mid-training. Every time\n  validation improves, this script already writes a full checkpoint (raptor_ft_<tag>.pt) that\n  contains the entire RaptorClassifier state_dict (backbone + attention head), not just the\n  backbone. Point --resume at that file to continue training on top of it in a new session:\n\n    python train_knee.py --resume /kaggle/working/raptor_ft_mymodel.pt --tag mymodel_cont \\\n        --arch coatnet_1_rw_224.sw_in1k --res 224 --epochs 6 --labels ...\n\n  --resume loads full model weights (backbone + head) and seeds `best` from the checkpoint's\n  saved gold_auc, so a fresh session won't clobber a better checkpoint with a worse one before\n  it catches back up. It does NOT restore optimizer or LR-scheduler state - training resumes\n  with a fresh OneCycleLR schedule over --epochs, so treat --resume as \"continue fine-tuning\n  from these weights,\" not \"restore an interrupted session bit-for-bit.\" --arch and --res\n  should match the checkpoint you're resuming from, or shapes won't load; a mismatch prints\n  a warning rather than failing outright, so check that warning if you see one.\n\"\"\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T14:31:05.741497Z","iopub.execute_input":"2026-08-26T14:31:05.74175Z","iopub.status.idle":"2026-08-26T14:31:05.752692Z","shell.execute_reply.started":"2026-08-26T14:31:05.741729Z","shell.execute_reply":"2026-08-26T14:31:05.751933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile train_knee.py\n\nimport os, sys, time, json, math, random, argparse\nimport numpy as np, pandas as pd\nimport torch, torch.nn as nn, torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nimport timm\nfrom tqdm.auto import tqdm\n\n\n# --- ROI localizer (anatomical joint crop). Optional so the no-ROI path is untouched. ---\nsys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))\ntry:\n    from roi_localize import square_box as _roi_square_box, compartments as _roi_compartments\nexcept Exception:\n    _roi_square_box = _roi_compartments = None\n\nHERE = os.path.dirname(os.path.abspath(__file__))\nRSNA = os.path.dirname(HERE)\n\n# ---------------------------------------------------------------------------\n# Input discovery. On Kaggle the corpus arrives as two read-only datasets and the\n# competition data as a third, so nothing lives beside this script. Everything in\n# this block is discovery only - the training code further down is unchanged.\n# ---------------------------------------------------------------------------\ndef _find(*names, root=\"/kaggle/input\"):\n    \"\"\"First path under root whose basename matches one of names.\"\"\"\n    for d, _, fs in os.walk(root):\n        for n in names:\n            if n in fs:\n                return os.path.join(d, n)\n    return None\n\n\nclass _TwoPartVols:\n    \"\"\"Presents the two published corpus parts as one array of shape (4407, 44, 336, 336).\n\n    Both parts stay memory-mapped and are never concatenated on disk: copying 22 GB would be\n    pointless when every read is a single study. Row order is part 1 then part 2, matching the\n    order the id files concatenate in. That ordering is the contract between volumes, masks and\n    ids, so do not sort any of them independently.\n    \"\"\"\n    def __init__(self, a, b):\n        self.a, self.b = a, b\n        self.n_a = a.shape[0]\n        self.shape = (a.shape[0] + b.shape[0],) + tuple(a.shape[1:])\n\n    def __len__(self):\n        return self.shape[0]\n\n    def __getitem__(self, row):\n        return self.a[row] if row < self.n_a else self.b[row - self.n_a]\n\n\ndef _open_corpus():\n    \"\"\"Return (vols, masks), from a local single-file corpus or the two public parts.\"\"\"\n    local_v = os.path.join(HERE, \"all_vols.npy\")\n    if os.path.exists(local_v):\n        return (np.load(local_v, mmap_mode=\"r\"),\n                np.load(os.path.join(HERE, \"all_masks.npy\")))\n    av, bv = _find(\"all_vols.npy\"), _find(\"extra_vols.npy\")\n    am, bm = _find(\"all_masks.npy\"), _find(\"extra_masks.npy\")\n    if not all((av, bv, am, bm)):\n        raise SystemExit(\"Could not find the corpus. Attach both parts: \"\n                         \"dreaddevelopment/knee-raptor-corpus and \"\n                         \"dreaddevelopment/knee-raptor-corpus-ext\")\n    vols = _TwoPartVols(np.load(av, mmap_mode=\"r\"), np.load(bv, mmap_mode=\"r\"))\n    masks = np.concatenate([np.load(am), np.load(bm)], axis=0)\n    return vols, masks\n\n\ndef _open_ids():\n    local = os.path.join(HERE, \"all_ids.npy\")\n    if os.path.exists(local):\n        return np.load(local, allow_pickle=True).astype(str)\n    a, b = _find(\"all_ids.npy\"), _find(\"extra_ids.npy\")\n    if not (a and b):\n        raise SystemExit(\"Could not find all_ids.npy / extra_ids.npy - attach both corpus parts.\")\n    return np.concatenate([np.load(a, allow_pickle=True).astype(str),\n                           np.load(b, allow_pickle=True).astype(str)])\nLAB = [\"ACL\",\"MCL\",\"Medial Meniscus\",\"Lateral Meniscus\",\"Medial OA\",\"Lateral OA\",\"PF OA\",\n       \"Effusion\",\"Synovitis\",\"Baker's\",\"Contusion\",\"Fracture\"]\n\n\n_NO_LABELS = \"\"\"\nNo training labels found, so there is nothing to train against.\n\nThe competition labels only 58 of the 4,407 studies. The other 4,349 carry a free-text\nradiology report instead, so before this can train you have to turn those reports into\ntwelve probabilities per study and pass the result with --labels.\n\nExpected format: a parquet with a StudyInstanceUID column plus the columns\n  {cols}\nwith values between 0 and 1. Soft values work better than hard 0/1 here: the loss is built\nfor them, and hedged reports are common.\n\nEverything else in this notebook is ready to run once that file exists.\n\"\"\"\n\n\n# ------------------------------- data ----------------------------------------\nclass StudyWindows(Dataset):\n    \"\"\"Per-study bag of 2.5D windows sampled from all_vols.npy (memmap).\n    Each window = 3 physically-consecutive slices -> RGB, resized to `res`, in [0,1]\n    (matches the SSL input pipeline: ToTensor, no ImageNet norm).\"\"\"\n    def __init__(self, root, ids, id2row, labels, res, k, train, aug=True, norm=\"none\",\n                 roi=False, roi_mode=\"tight\", roi_pad=0.06, roi_overlap=0.12,\n                 roi_sbox=None, roi_cen=None):\n        self.root = root\n        self.ids = ids\n        self.id2row = id2row\n        self.labels = labels              # dict uid -> np.float32[12]\n        self.res, self.k, self.train, self.aug = res, k, train, aug\n        self.norm = norm                  # \"none\"=[0,1] (Raptor SSL); \"imagenet\"=DINOv2 stats\n        self._mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n        self._std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n        self.vols = None; self.masks = None\n        # --- ROI anatomical joint-crop config ---\n        self.roi = bool(roi)\n        self.roi_mode, self.roi_pad, self.roi_overlap = roi_mode, roi_pad, roi_overlap\n        self.roi_sbox = roi_sbox          # (N,D,4) int16 per-slice tissue bbox, or None\n        self.roi_cen = roi_cen            # (N,D,2) int16 per-slice joint centroid, or None\n        if self.roi and (_roi_square_box is None or roi_sbox is None):\n            raise RuntimeError(\"roi=True but roi_localize or roi_boxes not available\")\n\n    def __len__(self): return len(self.ids)\n\n    def _ensure(self):\n        if self.vols is None:\n            self.vols, self.masks = _open_corpus()   # (N,D,H,W) view, (N,D) u8\n\n    def _centers(self, valid, count):\n        # valid slice indices; window centers must have both neighbors valid & in-range\n        lo, hi = int(valid.min()), int(valid.max())\n        cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n        if not cs: cs = [max(1, min((lo + hi) // 2, self._D - 2))]\n        if self.train:\n            reps = count // len(cs) + 1\n            pool = (cs * reps)\n            random.shuffle(pool)\n            return pool[:count]\n        # eval: evenly spaced deterministic\n        idx = np.linspace(0, len(cs) - 1, count).round().astype(int)\n        return [cs[i] for i in idx]\n\n    def _resize(self, tri):\n        \"\"\"tri: (3,h,w) float32 [0,1] -> (3,res,res) float32.\"\"\"\n        t = torch.from_numpy(np.ascontiguousarray(tri))\n        if t.shape[-1] != self.res or t.shape[-2] != self.res:\n            t = F.interpolate(t[None], size=(self.res, self.res), mode=\"bilinear\",\n                              align_corners=False)[0]\n        return t.numpy()\n\n    def __getitem__(self, i):\n        self._ensure()\n        uid = self.ids[i]; row = self.id2row[uid]\n        self._D = self.vols.shape[1]\n        m = self.masks[row]\n        valid = np.where(m > 0)[0]\n        if len(valid) < 3: valid = np.arange(min(3, self._D))\n        # compartment mode emits 2 crops/center -> sample ceil(k/2) centers to keep #windows==k\n        compart = self.roi and self.roi_mode == \"compartment\"\n        n_centers = (self.k + 1) // 2 if compart else self.k\n        cs = self._centers(valid, n_centers)\n        vol = self.vols[row]  # (D,H,W) u8  (single study read)\n        tiles = []            # list of (3,res,res) float32\n        for c in cs:\n            c = max(1, min(c, self._D - 2))\n            tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0  # (3,H,W)\n            H, W = tri.shape[-2], tri.shape[-1]\n            if not self.roi:\n                tiles.append(self._resize(tri))\n                continue\n            # one joint box from the CENTER slice, applied to all 3 slices (keeps RGB registered)\n            sq_mode = \"tight\" if compart else self.roi_mode\n            sq = _roi_square_box(tuple(int(v) for v in self.roi_sbox[row, c]), W=W, H=H,\n                                 pad=self.roi_pad, mode=sq_mode,\n                                 centroid=tuple(float(v) for v in self.roi_cen[row, c]))\n            if compart:\n                left, right = _roi_compartments(sq, overlap=self.roi_overlap)\n                for box in (left, right):\n                    x0, y0, x1, y1 = box\n                    tiles.append(self._resize(tri[:, y0:y1, x0:x1]))\n            else:\n                x0, y0, x1, y1 = sq\n                tiles.append(self._resize(tri[:, y0:y1, x0:x1]))\n        if len(tiles) > self.k:\n            tiles = tiles[:self.k]\n        wins = np.stack(tiles, 0)          # (K,3,res,res)\n        x = torch.from_numpy(wins)\n        if self.train and self.aug:\n            # light medical-safe aug: NO flips (laterality is signal); mild intensity jitter\n            g = 1.0 + (random.random() - 0.5) * 0.20\n            x = (x * g).clamp(0, 1)\n        if self.norm == \"imagenet\":        # each backbone at its correct input distribution\n            x = (x - self._mean) / self._std\n        y = torch.from_numpy(self.labels[uid])\n        return x, y\n\n\ndef collate(batch):\n    xs = torch.stack([b[0] for b in batch])   # (B,K,3,res,res)\n    ys = torch.stack([b[1] for b in batch])   # (B,12)\n    return xs, ys\n\n\n# ------------------------------- model ---------------------------------------\ndef build_backbone(arch=\"vit_small_patch16_224\", pretrained=False):\n    hybrid = arch.startswith((\"maxvit\", \"maxxvit\", \"coatnet\", \"coat_\", \"convnext\"))\n    is_vit = (not hybrid) and any(k in arch for k in (\"vit\", \"deit\", \"dinov2\", \"eva\", \"beit\"))\n    kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n    if is_vit:\n        kw.update(global_pool=\"token\", dynamic_img_size=True)\n    else:\n        kw.update(global_pool=\"avg\")\n    return timm.create_model(arch, **kw)\n\n\ndef load_raptor(bb, ckpt_path):\n    if ckpt_path in (\"timm\", \"pretrained\"):\n        return \"timm-pretrained\"\n    if ckpt_path in (\"\", \"none\", \"None\"):\n        print(\"[raptor] RANDOM-INIT control (no SSL weights)\", flush=True)\n        return \"random-init\"\n    ck = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n    st = ck[\"student\"] if \"student\" in ck else ck\n    bbst = {k[len(\"backbone.\"):]: v for k, v in st.items() if k.startswith(\"backbone.\")}\n    missing, unexpected = bb.load_state_dict(bbst, strict=False)\n    ep = ck.get(\"epoch\", \"?\")\n    print(f\"[raptor] loaded backbone from {os.path.basename(ckpt_path)} (ssl epoch {ep}) | \"\n          f\"loaded {len(bbst)} tensors, missing {len(missing)}, unexpected {len(unexpected)}\", flush=True)\n    return f\"{os.path.basename(ckpt_path)}@ep{ep}\"\n\n\nclass RaptorClassifier(nn.Module):\n    \"\"\"Raptor encoder + cross-window self-attention + per-diagnosis attention-MIL head.\"\"\"\n    def __init__(self, backbone, F_dim=384, n=12, drop=0.2, n_heads=4):\n        super().__init__()\n        self.backbone = backbone\n        # --- NEW: single cross-window self-attention layer (windows condition on each other) ---\n        self.mix_norm = nn.LayerNorm(F_dim)\n        self.mix = nn.MultiheadAttention(F_dim, n_heads, dropout=drop, batch_first=True)\n        # --- existing attention-MIL head, unchanged ---\n        self.norm = nn.LayerNorm(F_dim)\n        self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop),\n                                 nn.Linear(256, n))\n        self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n        self.clsb = nn.Parameter(torch.zeros(n))\n        nn.init.trunc_normal_(self.clsW, std=0.02)\n        self.n = n\n\n    def encode(self, x):\n        B, K = x.shape[:2]\n        f = self.backbone(x.flatten(0, 1))\n        return f.view(B, K, -1)\n\n    def mix_windows(self, feats):\n        # feats: (B,K,F). Self-attention across the K windows, pre-norm + residual.\n        h = self.mix_norm(feats)\n        h, _ = self.mix(h, h, h, need_weights=False)\n        return feats + h    # residual: transformer can learn identity if it's not helping\n\n    def head(self, feats):\n        h = self.norm(feats)\n        a = self.att(h)\n        a = torch.softmax(a, dim=1)\n        pooled = torch.einsum(\"bkn,bkf->bnf\", a, h)\n        logits = (pooled * self.clsW).sum(-1) + self.clsb\n        return logits\n\n    def forward(self, x):\n        feats = self.encode(x)\n        feats = self.mix_windows(feats)   # <-- the only new step\n        return self.head(feats)\n\n\n# ------------------------------- train ---------------------------------------\ndef main():\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--ckpt\", default=os.path.join(HERE, \"ckpt\", \"raptor_ssl_last.pt\"))\n    ap.add_argument(\"--arch\", default=\"vit_small_patch16_224\")\n    ap.add_argument(\"--res\", type=int, default=224)\n    ap.add_argument(\"--k\", type=int, default=12)\n    ap.add_argument(\"--k_eval\", type=int, default=24)\n    ap.add_argument(\"--epochs\", type=int, default=12)\n    ap.add_argument(\"--bs\", type=int, default=8)\n    ap.add_argument(\"--bb_lr\", type=float, default=1e-5)\n    ap.add_argument(\"--head_lr\", type=float, default=3e-4)\n    ap.add_argument(\"--wd\", type=float, default=0.02)\n    ap.add_argument(\"--workers\", type=int, default=4)\n    ap.add_argument(\"--limit\", type=int, default=0)\n    ap.add_argument(\"--freeze_blocks\", type=int, default=0)\n    ap.add_argument(\"--norm\", default=\"none\", choices=[\"none\", \"imagenet\"])\n    ap.add_argument(\"--grad_ckpt\", action=\"store_true\")\n    ap.add_argument(\"--tag\", default=\"dev\")\n    ap.add_argument(\"--labels\", default=None,\n                    help=\"parquet of training labels: StudyInstanceUID + the twelve finding \"\n                         \"columns, values 0..1. Not provided with this notebook - see the header.\")\n    ap.add_argument(\"--smoke\", action=\"store_true\")\n    ap.add_argument(\"--seed\", type=int, default=42)\n    # ---- CV fold hook ----\n    ap.add_argument(\"--folds\", type=int, default=0)\n    ap.add_argument(\"--fold\", type=int, default=-1)\n    ap.add_argument(\"--fold_file\", default=None)\n    # ---- anatomical ROI joint-crop (A/B lever). Default OFF -> identical to baseline. ----\n    ap.add_argument(\"--roi\", action=\"store_true\")\n    ap.add_argument(\"--roi_mode\", default=\"compartment\", choices=[\"tight\", \"safe\", \"compartment\"])\n    ap.add_argument(\"--roi_pad\", type=float, default=0.06)\n    ap.add_argument(\"--roi_overlap\", type=float, default=0.12)\n    ap.add_argument(\"--roi_boxes\", default=os.path.join(HERE, \"roi_boxes.npz\"))\n    # ---- resume across sessions: continue on top of a previously saved raptor_ft_*.pt ----\n    ap.add_argument(\"--resume\", default=None,\n                    help=\"path to a previous raptor_ft_<tag>.pt (full RaptorClassifier \"\n                         \"state_dict) to resume training from. Loads model weights and seeds \"\n                         \"the best-so-far gold_auc so a fresh session won't overwrite a better \"\n                         \"checkpoint before catching up. Does NOT restore optimizer/scheduler \"\n                         \"state - training restarts with a fresh OneCycleLR over --epochs.\")\n    a = ap.parse_args()\n    random.seed(a.seed); np.random.seed(a.seed); torch.manual_seed(a.seed)\n    dev = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    if a.smoke:\n        a.epochs, a.bs, a.k, a.k_eval, a.limit, a.workers = 2, 4, 4, 6, 40, 0\n    print(f\"device {dev} | res {a.res} | k {a.k}/{a.k_eval} | bs {a.bs} | tag {a.tag} | \"\n          f\"roi={a.roi}({a.roi_mode})\", flush=True)\n\n    # ---- ids / labels ----\n    ids = _open_ids()\n    id2row = {u: i for i, u in enumerate(ids)}\n    idset = set(ids)\n    _tcsv = os.path.join(RSNA, \"train.csv\")\n    if not os.path.exists(_tcsv):\n        _tcsv = _find(\"train.csv\")\n    if not _tcsv:\n        raise SystemExit(\"Could not find train.csv - attach the competition data.\")\n    tr = pd.read_csv(_tcsv); tr[\"StudyInstanceUID\"] = tr[\"StudyInstanceUID\"].astype(str)\n    gold_df = tr[tr[LAB].notna().all(axis=1)].copy().set_index(\"StudyInstanceUID\")\n    gold_ids = [u for u in gold_df.index if u in idset]\n    _lab = a.labels or _find(\"labels_llm_soft.parquet\")\n    if not _lab or not os.path.exists(_lab):\n        # Exit cleanly rather than as a failure: running this notebook as published, with no\n        # labels attached, is the expected path and should read as an explanation, not a crash.\n        print(_NO_LABELS.format(cols=\", \".join(LAB)), flush=True)\n        raise SystemExit(0)\n    soft = pd.read_parquet(_lab)\n    soft[\"StudyInstanceUID\"] = soft[\"StudyInstanceUID\"].astype(str); soft = soft.set_index(\"StudyInstanceUID\")\n    goldset = set(gold_ids)\n    train_ids = [u for u in ids if u in soft.index and u not in goldset]\n    # ---- CV fold hook: hold out fold `a.fold`, train on the rest ----\n    oof_ids = []\n    if a.folds > 0:\n        assert 0 <= a.fold < a.folds, f\"--fold must be in [0,{a.folds}) when --folds>0\"\n        fmap = json.load(open(a.fold_file))[\"folds\"]\n        held = set(u for u in train_ids if fmap.get(u, -1) == a.fold)\n        oof_ids = [u for u in train_ids if u in held]\n        train_ids = [u for u in train_ids if u not in held]\n        print(f\"[cv] fold {a.fold}/{a.folds}: train {len(train_ids)} | OOF held-out {len(oof_ids)} \"\n              f\"| fold_file {os.path.basename(a.fold_file)}\", flush=True)\n    if a.limit: train_ids = train_ids[:a.limit]\n    labels = {u: soft.loc[u, LAB].values.astype(np.float32) for u in train_ids}\n    for u in oof_ids: labels[u] = soft.loc[u, LAB].values.astype(np.float32)\n    for u in gold_ids: labels[u] = gold_df.loc[u, LAB].values.astype(np.float32)\n    print(f\"train {len(train_ids)} | gold-val {len(gold_ids)}\"\n          + (f\" | oof {len(oof_ids)}\" if oof_ids else \"\"), flush=True)\n\n    prev = np.clip(np.stack([labels[u] for u in train_ids]).mean(0), 0.03, 0.7)\n    pw = torch.tensor(np.clip((1 - prev) / prev, 1, 10), dtype=torch.float32, device=dev)\n\n    # ---- ROI boxes (only when --roi) ----\n    roi_sbox = roi_cen = None\n    if a.roi:\n        rb = np.load(a.roi_boxes)\n        rb_ids = rb[\"ids\"].astype(str)\n        if not np.array_equal(rb_ids, ids):\n            rmap = {u: i for i, u in enumerate(rb_ids)}\n            missing = [u for u in ids if u not in rmap]\n            if missing:\n                raise RuntimeError(f\"roi_boxes missing {len(missing)} corpus ids (e.g. {missing[:2]})\")\n            order = np.array([rmap[u] for u in ids])\n            roi_sbox = rb[\"sbox\"][order]; roi_cen = rb[\"cen\"][order]\n        else:\n            roi_sbox = rb[\"sbox\"]; roi_cen = rb[\"cen\"]\n        print(f\"[roi] ENABLED mode={a.roi_mode} pad={a.roi_pad} overlap={a.roi_overlap} \"\n              f\"boxes={os.path.basename(a.roi_boxes)} sbox={roi_sbox.shape}\", flush=True)\n    _roi_kw = dict(roi=a.roi, roi_mode=a.roi_mode, roi_pad=a.roi_pad, roi_overlap=a.roi_overlap,\n                   roi_sbox=roi_sbox, roi_cen=roi_cen)\n\n    tds = StudyWindows(HERE, train_ids, id2row, labels, a.res, a.k, train=True, norm=a.norm, **_roi_kw)\n    vds = StudyWindows(HERE, gold_ids, id2row, labels, a.res, a.k_eval, train=False, norm=a.norm, **_roi_kw)\n    tl = DataLoader(tds, batch_size=a.bs, shuffle=True, num_workers=a.workers, drop_last=True,\n                    collate_fn=collate, pin_memory=True, persistent_workers=a.workers > 0)\n    vl = DataLoader(vds, batch_size=max(2, a.bs // 2), shuffle=False, num_workers=a.workers,\n                    collate_fn=collate, persistent_workers=a.workers > 0)\n\n    use_timm = a.ckpt in (\"timm\", \"pretrained\")\n    bb = build_backbone(a.arch, pretrained=use_timm)\n    src = load_raptor(bb, a.ckpt)\n    if use_timm: print(f\"[raptor] timm-pretrained backbone: {a.arch}\", flush=True)\n    F_dim = bb.num_features\n    if a.grad_ckpt:\n        try:\n            bb.set_grad_checkpointing(True); print(\"[raptor] gradient checkpointing ON\", flush=True)\n        except Exception as e:\n            print(f\"[raptor] grad_ckpt unsupported for {a.arch}: {e}\", flush=True)\n    model = RaptorClassifier(bb, F_dim=F_dim).to(dev)\n    if a.freeze_blocks > 0:\n        for nm, p in model.backbone.named_parameters():\n            for b in range(a.freeze_blocks):\n                if nm.startswith(f\"blocks.{b}.\"): p.requires_grad = False\n\n    # ---- resume: load full model weights (backbone + attention head) from a prior run ----\n    resume_best = 0.0\n    if a.resume:\n        rck = torch.load(a.resume, map_location=\"cpu\", weights_only=False)\n        r_arch, r_res = rck.get(\"arch\"), rck.get(\"res\")\n        if r_arch is not None and r_arch != a.arch:\n            print(f\"[resume] WARNING: checkpoint arch '{r_arch}' != current --arch '{a.arch}' \"\n                  f\"- shapes will likely mismatch\", flush=True)\n        if r_res is not None and r_res != a.res:\n            print(f\"[resume] WARNING: checkpoint res {r_res} != current --res {a.res} \"\n                  f\"- fine for dynamic-size ViTs, but double check for others\", flush=True)\n        missing, unexpected = model.load_state_dict(rck[\"model\"], strict=False)\n        resume_best = float(rck.get(\"gold_auc\") or 0.0)\n        print(f\"[resume] loaded {os.path.basename(a.resume)} | prior gold_auc {resume_best:.4f} | \"\n              f\"missing {len(missing)} unexpected {len(unexpected)}\", flush=True)\n        if missing or unexpected:\n            print(f\"[resume]   missing keys (first 5): {missing[:5]}\", flush=True)\n            print(f\"[resume]   unexpected keys (first 5): {unexpected[:5]}\", flush=True)\n\n    head_params = [p for n_, p in model.named_parameters() if not n_.startswith(\"backbone.\") and p.requires_grad]\n    bb_params = [p for n_, p in model.named_parameters() if n_.startswith(\"backbone.\") and p.requires_grad]\n    opt = torch.optim.AdamW([{\"params\": bb_params, \"lr\": a.bb_lr},\n                             {\"params\": head_params, \"lr\": a.head_lr}], weight_decay=a.wd)\n    steps = max(len(tl) * a.epochs, 1)\n    sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=[a.bb_lr, a.head_lr],\n                                                total_steps=steps, pct_start=0.15)\n    lossf = nn.BCEWithLogitsLoss(pos_weight=pw)\n\n    @torch.no_grad()\n    def evaluate():\n        model.eval(); P = []; Y = []\n        for x, y in vl:\n            with torch.autocast(dev, dtype=torch.bfloat16, enabled=dev == \"cuda\"):\n                o = torch.sigmoid(model(x.to(dev)).float())\n            P.append(o.cpu().numpy()); Y.append(y.numpy())\n        P = np.concatenate(P); Y = np.concatenate(Y)\n        aucs = {}\n        for j, name in enumerate(LAB):\n            if len(set(Y[:, j].astype(int))) > 1:\n                aucs[name] = float(roc_auc_score(Y[:, j], P[:, j]))\n        return float(np.mean(list(aucs.values()))), aucs, P, Y\n\n    best = resume_best; best_state = None; best_P = None; t0 = time.time(); hist = []\n    TOPK = 3; topk = []\n    for ep in range(a.epochs):\n        model.train(); tot = 0.0\n        pbar = tqdm(tl, desc=f\"ep{ep}\", leave=False)\n        for step, (x, y) in enumerate(pbar):\n            x, y = x.to(dev, non_blocking=True), y.to(dev, non_blocking=True)\n            opt.zero_grad()\n            with torch.autocast(dev, dtype=torch.bfloat16, enabled=dev == \"cuda\"):\n                logits = model(x)\n                loss = lossf(logits.float(), y)\n            if not torch.isfinite(loss):\n                print(f\"🚨 NaN/Inf loss at epoch={ep}, step={step}\")\n                print(f\"loss = {loss.item()}\")\n                print(f\"x finite = {torch.isfinite(x).all().item()}\")\n                print(f\"y finite = {torch.isfinite(y).all().item()}\")\n                print(f\"logits finite = {torch.isfinite(logits).all().item()}\")\n                raise RuntimeError(\"NaN/Inf loss\")\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 3.0)\n            opt.step(); sched.step(); tot += loss.item()\n            for name, p in model.named_parameters():\n                if p.requires_grad and not torch.isfinite(p).all():\n                    print(f\"🚨 NaN/Inf parameter at epoch={ep}, step={step}: {name}\")\n                    raise RuntimeError(f\"Model parameter became NaN/Inf: {name}\")\n            pbar.set_postfix(loss=f\"{tot/(step+1):.3f}\")\n            \n        au, aucs, P, Y = evaluate()\n        hist.append({\"ep\": ep, \"loss\": tot / len(tl), \"gold_auc\": au})\n        if au > best:\n            best = au; best_P = P\n            best_state = {\"model\": {k: v.detach().cpu() for k, v in model.state_dict().items()},\n                          \"gold_auc\": au, \"aucs\": aucs, \"src\": src, \"res\": a.res,\n                          \"arch\": a.arch, \"lab\": LAB, \"epoch\": ep}\n        if au >= best and best_state is not None:\n            _tmp = os.path.join(HERE, f\"raptor_ft_{a.tag}.pt.tmp\")\n            torch.save(best_state, _tmp)\n            os.replace(_tmp, os.path.join(HERE, f\"raptor_ft_{a.tag}.pt\"))\n            np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}.npz\"),\n                     pred=best_P, truth=Y, ids=np.array(gold_ids))\n            json.dump({\"tag\": a.tag, \"src\": src, \"best_gold_auc\": best,\n                       \"aucs\": best_state[\"aucs\"], \"hist\": hist, \"res\": a.res,\n                       \"epochs\": a.epochs, \"bb_lr\": a.bb_lr, \"head_lr\": a.head_lr,\n                       \"n_train\": len(train_ids), \"n_gold\": len(gold_ids),\n                       \"roi\": a.roi, \"roi_mode\": a.roi_mode,\n                       \"resumed_from\": a.resume,\n                       \"partial\": True, \"epochs_done\": ep + 1},\n                      open(os.path.join(HERE, f\"raptor_ft_{a.tag}.json\"), \"w\"), indent=1)\n            print(f\"  [ckpt] best-so-far saved at ep{ep} ({best:.4f})\", flush=True)\n        if len(topk) < TOPK or au > min(t[\"gold_auc\"] for t in topk):\n            topk.append({\"model\": {k: v.detach().cpu().clone() for k, v in model.state_dict().items()},\n                         \"gold_auc\": au, \"aucs\": aucs, \"src\": src, \"res\": a.res,\n                         \"arch\": a.arch, \"lab\": LAB, \"epoch\": ep, \"P\": P})\n            topk.sort(key=lambda t: -t[\"gold_auc\"])\n            del topk[TOPK:]\n        print(f\"ep{ep} loss {tot/len(tl):.3f} | GOLD macro-AUC {au:.4f} (best {best:.4f}) | {time.time()-t0:.0f}s\",\n              flush=True)\n    _, aucs, _, Y = evaluate()\n    print(f\"\\nDONE {src} | BEST GOLD macro-AUC {best:.4f}\", flush=True)\n    for k, v in (best_state[\"aucs\"] if best_state else aucs).items():\n        print(f\"   {k:18s} {v:.3f}\", flush=True)\n\n    if best_state is not None:\n        torch.save(best_state, os.path.join(HERE, f\"raptor_ft_{a.tag}.pt\"))\n        np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}.npz\"),\n                 pred=best_P, truth=Y, ids=np.array(gold_ids))\n        json.dump({\"tag\": a.tag, \"src\": src, \"best_gold_auc\": best, \"aucs\": best_state[\"aucs\"],\n                   \"hist\": hist, \"res\": a.res, \"epochs\": a.epochs, \"bb_lr\": a.bb_lr,\n                   \"head_lr\": a.head_lr, \"n_train\": len(train_ids), \"n_gold\": len(gold_ids),\n                   \"roi\": a.roi, \"roi_mode\": a.roi_mode, \"resumed_from\": a.resume,\n                   \"partial\": False, \"epochs_done\": a.epochs},\n                  open(os.path.join(HERE, f\"raptor_ft_{a.tag}.json\"), \"w\"), indent=1)\n        print(f\"saved raptor_ft_{a.tag}.pt / raptor_gold_{a.tag}.npz / raptor_ft_{a.tag}.json\", flush=True)\n\n    # ---- CV OOF: predict the held-out fold at the best (gold-selected) weights ----\n    if a.folds > 0 and oof_ids and best_state is not None:\n        model.load_state_dict(best_state[\"model\"]); model.eval()\n        ods = StudyWindows(HERE, oof_ids, id2row, labels, a.res, a.k_eval, train=False, norm=a.norm, **_roi_kw)\n        ol = DataLoader(ods, batch_size=max(2, a.bs // 2), shuffle=False, num_workers=a.workers,\n                        collate_fn=collate, persistent_workers=False)\n        Po = []\n        with torch.no_grad():\n            for x, y in ol:\n                with torch.autocast(dev, dtype=torch.bfloat16, enabled=dev == \"cuda\"):\n                    o = torch.sigmoid(model(x.to(dev)).float())\n                Po.append(o.cpu().numpy())\n        Po = np.concatenate(Po)\n        Yo = np.stack([labels[u] for u in oof_ids]).astype(np.float32)\n        np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_fold{a.fold}.npz\"),\n                 pred=Po, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n        print(f\"[cv] wrote raptor_oof_{a.tag}_fold{a.fold}.npz ({len(oof_ids)} studies; \"\n              f\"truth = soft labels)\", flush=True)\n\n        # --- top-K epoch OOF (new) ---\n        # The epoch-ensemble used to be judged on gold only; that gate is too small to\n        # resolve the move. Re-run the held-out fold at each retained epoch instead.\n        oof_by_ep = {best_state[\"epoch\"]: Po}\n        for t in topk:\n            e = int(t[\"epoch\"])\n            if e in oof_by_ep:\n                continue\n            model.load_state_dict(t[\"model\"]); model.eval()\n            Pe = []\n            with torch.no_grad():\n                for x, y in ol:\n                    with torch.autocast(dev, dtype=torch.bfloat16, enabled=dev == \"cuda\"):\n                        Pe.append(torch.sigmoid(model(x.to(dev)).float()).cpu().numpy())\n            Pe = np.concatenate(Pe); oof_by_ep[e] = Pe\n            np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_ep{e}_fold{a.fold}.npz\"),\n                     pred=Pe, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n            print(f\"[cv] wrote top-K epoch OOF ep{e}\", flush=True)\n        if len(oof_by_ep) > 1:\n            Pens = np.mean(list(oof_by_ep.values()), axis=0)\n            np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_epens_fold{a.fold}.npz\"),\n                     pred=Pens, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n            print(f\"[cv] wrote epoch-ensemble OOF over epochs {sorted(oof_by_ep)}\", flush=True)\n        model.load_state_dict(best_state[\"model\"]); model.eval()\n\n    # --- weight-averaged checkpoint (new) ---\n    if len(topk) > 1:\n        import copy\n        sds = [t[\"model\"] for t in topk]\n        avg = {}\n        for k in sds[0]:\n            v0 = sds[0][k]\n            if v0.is_floating_point():\n                avg[k] = sum(sd[k].double() for sd in sds).div(len(sds)).to(v0.dtype)\n            else:\n                avg[k] = v0.clone()          # e.g. num_batches_tracked\n        swa_state = {\"model\": avg, \"gold_auc\": None, \"aucs\": {}, \"src\": src, \"res\": a.res,\n                     \"arch\": a.arch, \"lab\": LAB, \"epoch\": [int(t[\"epoch\"]) for t in topk],\n                     \"swa_over\": [int(t[\"epoch\"]) for t in topk]}\n        model.load_state_dict(avg); model.eval()\n        au_swa, aucs_swa, P_swa, Y_swa = evaluate()\n        swa_state[\"gold_auc\"] = au_swa; swa_state[\"aucs\"] = aucs_swa\n        torch.save(swa_state, os.path.join(HERE, f\"raptor_ft_{a.tag}_swa.pt\"))\n        np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_swa.npz\"),\n                 pred=P_swa, truth=Y_swa, ids=np.array(gold_ids))\n        print(f\"SWA over epochs {[int(t['epoch']) for t in topk]} | gold {au_swa:.4f} \"\n              f\"(best-epoch {best:.4f}) {'BETTER' if au_swa > best else 'no gain'}\", flush=True)\n        if a.folds > 0 and oof_ids:\n            ods2 = StudyWindows(HERE, oof_ids, id2row, labels, a.res, a.k_eval, train=False,\n                                norm=a.norm, **_roi_kw)\n            ol2 = DataLoader(ods2, batch_size=max(2, a.bs // 2), shuffle=False,\n                             num_workers=a.workers, collate_fn=collate, persistent_workers=False)\n            Ps = []\n            with torch.no_grad():\n                for x, y in ol2:\n                    with torch.autocast(dev, dtype=torch.bfloat16, enabled=dev == \"cuda\"):\n                        Ps.append(torch.sigmoid(model(x.to(dev)).float()).cpu().numpy())\n            Ps = np.concatenate(Ps)\n            np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_swa_fold{a.fold}.npz\"),\n                     pred=Ps, truth=np.stack([labels[u] for u in oof_ids]).astype(np.float32),\n                     ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n            print(f\"[cv] wrote SWA OOF\", flush=True)\n        model.load_state_dict(best_state[\"model\"]); model.eval()\n\n    if len(topk) > 1:\n        eps = [t[\"epoch\"] for t in topk]\n        for rank, t in enumerate(topk):\n            if rank == 0:\n                continue\n            P_t = t.pop(\"P\")\n            np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_ep{t['epoch']}.npz\"),\n                     pred=P_t, truth=Y, ids=np.array(gold_ids))\n            t[\"P\"] = P_t\n        Pens = np.mean([t[\"P\"] for t in topk], axis=0)\n        ens_aucs = {}\n        for j, name in enumerate(LAB):\n            if len(set(Y[:, j].astype(int))) > 1:\n                ens_aucs[name] = float(roc_auc_score(Y[:, j], Pens[:, j]))\n        ens = float(np.mean(list(ens_aucs.values())))\n        np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_epens.npz\"),\n                 pred=Pens, truth=Y, ids=np.array(gold_ids))\n        print(f\"TOPK epochs {eps} | best {best:.4f} | epoch-ensemble {ens:.4f} \"\n              f\"({'BETTER' if ens > best else 'no gain'})\", flush=True)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:26:47.289851Z","iopub.execute_input":"2026-08-26T12:26:47.290148Z","iopub.status.idle":"2026-08-26T12:26:47.313853Z","shell.execute_reply.started":"2026-08-26T12:26:47.290125Z","shell.execute_reply":"2026-08-26T12:26:47.31329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python train_knee.py \\\n    --arch coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k \\\n    --res 384 \\\n    --epochs 1 \\\n    --bs 4 \\\n    --k 12 \\\n    --k_eval 24 \\\n    --grad_ckpt \\\n    --workers 2 \\\n    --tag debug200 \\\n    --ckpt timm \\\n    --bb_lr 1e-5 \\\n    --head_lr 3e-4 \\\n    --limit 1000 \\\n    --labels /kaggle/working/labels_soft.parquet \\\n    --resume /kaggle/input/models/yugeshkctheaiman/test/pytorch/default/1/raptor_ft_debug200.pt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:52.543783Z","iopub.execute_input":"2026-08-26T12:52:52.544053Z","iopub.status.idle":"2026-08-26T12:52:55.927448Z","shell.execute_reply.started":"2026-08-26T12:52:52.544024Z","shell.execute_reply":"2026-08-26T12:52:55.926497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python train_knee.py --arch coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k --res 384 \\\n#     --epochs 2 --bs 4 --k 12 --k_eval 24 --grad_ckpt --workers 2 \\\n#     --tag firstok --ckpt timm \\\n#     --labels /kaggle/working/labels_soft.parquet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.688459Z","iopub.execute_input":"2026-08-26T12:52:46.688786Z","iopub.status.idle":"2026-08-26T12:52:46.692793Z","shell.execute_reply.started":"2026-08-26T12:52:46.688744Z","shell.execute_reply":"2026-08-26T12:52:46.692121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !grep -n \"tqdm\" /kaggle/working/train_knee.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.693785Z","iopub.execute_input":"2026-08-26T12:52:46.694451Z","iopub.status.idle":"2026-08-26T12:52:46.705784Z","shell.execute_reply.started":"2026-08-26T12:52:46.694388Z","shell.execute_reply":"2026-08-26T12:52:46.705029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !ps aux | grep train_knee","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.706681Z","iopub.execute_input":"2026-08-26T12:52:46.706955Z","iopub.status.idle":"2026-08-26T12:52:46.716112Z","shell.execute_reply.started":"2026-08-26T12:52:46.706923Z","shell.execute_reply":"2026-08-26T12:52:46.715472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !grep -n \"resume\" /kaggle/working/train_knee.py | head -20","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.716865Z","iopub.execute_input":"2026-08-26T12:52:46.717145Z","iopub.status.idle":"2026-08-26T12:52:46.729365Z","shell.execute_reply.started":"2026-08-26T12:52:46.717117Z","shell.execute_reply":"2026-08-26T12:52:46.72877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile train_knee.py\n# import os, sys, time, json, math, random, argparse\n# import numpy as np, pandas as pd\n# import torch, torch.nn as nn, torch.nn.functional as F\n# from torch.utils.data import Dataset, DataLoader\n# from sklearn.metrics import roc_auc_score\n# import timm\n\n# # --- ROI localizer (anatomical joint crop). Optional so the no-ROI path is untouched. ---\n# sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))\n# try:\n#     from roi_localize import square_box as _roi_square_box, compartments as _roi_compartments\n# except Exception:\n#     _roi_square_box = _roi_compartments = None\n\n# HERE = os.path.dirname(os.path.abspath(__file__))\n# RSNA = os.path.dirname(HERE)\n# print(\"here\",HERE)\n\n# # ---------------------------------------------------------------------------\n# # Input discovery. On Kaggle the corpus arrives as two read-only datasets and the\n# # competition data as a third, so nothing lives beside this script. Everything in\n# # this block is discovery only - the training code further down is unchanged.\n# # ---------------------------------------------------------------------------\n# def _find(*names, root=\"/kaggle/input\"):\n#     \"\"\"First path under root whose basename matches one of names.\"\"\"\n#     for d, _, fs in os.walk(root):\n#         for n in names:\n#             if n in fs:\n#                 return os.path.join(d, n)\n#     return None\n\n\n# class _TwoPartVols:\n#     \"\"\"Presents the two published corpus parts as one array of shape (4407, 44, 336, 336).\n\n#     Both parts stay memory-mapped and are never concatenated on disk: copying 22 GB would be\n#     pointless when every read is a single study. Row order is part 1 then part 2, matching the\n#     order the id files concatenate in. That ordering is the contract between volumes, masks and\n#     ids, so do not sort any of them independently.\n#     \"\"\"\n#     def __init__(self, a, b):\n#         self.a, self.b = a, b\n#         self.n_a = a.shape[0]\n#         self.shape = (a.shape[0] + b.shape[0],) + tuple(a.shape[1:])\n\n#     def __len__(self):\n#         return self.shape[0]\n\n#     def __getitem__(self, row):\n#         return self.a[row] if row < self.n_a else self.b[row - self.n_a]\n\n\n# def _open_corpus():\n#     \"\"\"Return (vols, masks), from a local single-file corpus or the two public parts.\"\"\"\n#     local_v = os.path.join(HERE, \"all_vols.npy\")\n#     if os.path.exists(local_v):\n#         return (np.load(local_v, mmap_mode=\"r\"),\n#                 np.load(os.path.join(HERE, \"all_masks.npy\")))\n#     av, bv = _find(\"all_vols.npy\"), _find(\"extra_vols.npy\")\n#     am, bm = _find(\"all_masks.npy\"), _find(\"extra_masks.npy\")\n#     if not all((av, bv, am, bm)):\n#         raise SystemExit(\"Could not find the corpus. Attach both parts: \"\n#                          \"dreaddevelopment/knee-raptor-corpus and \"\n#                          \"dreaddevelopment/knee-raptor-corpus-ext\")\n#     vols = _TwoPartVols(np.load(av, mmap_mode=\"r\"), np.load(bv, mmap_mode=\"r\"))\n#     masks = np.concatenate([np.load(am), np.load(bm)], axis=0)\n#     return vols, masks\n\n\n# def _open_ids():\n#     local = os.path.join(HERE, \"all_ids.npy\")\n#     if os.path.exists(local):\n#         return np.load(local, allow_pickle=True).astype(str)\n#     a, b = _find(\"all_ids.npy\"), _find(\"extra_ids.npy\")\n#     if not (a and b):\n#         raise SystemExit(\"Could not find all_ids.npy / extra_ids.npy - attach both corpus parts.\")\n#     return np.concatenate([np.load(a, allow_pickle=True).astype(str),\n#                            np.load(b, allow_pickle=True).astype(str)])\n# LAB = [\"ACL\",\"MCL\",\"Medial Meniscus\",\"Lateral Meniscus\",\"Medial OA\",\"Lateral OA\",\"PF OA\",\n#        \"Effusion\",\"Synovitis\",\"Baker's\",\"Contusion\",\"Fracture\"]\n\n\n# _NO_LABELS = \"\"\"\n# No training labels found, so there is nothing to train against.\n\n# The competition labels only 58 of the 4,407 studies. The other 4,349 carry a free-text\n# radiology report instead, so before this can train you have to turn those reports into\n# twelve probabilities per study and pass the result with --labels.\n\n# Expected format: a parquet with a StudyInstanceUID column plus the columns\n#   {cols}\n# with values between 0 and 1. Soft values work better than hard 0/1 here: the loss is built\n# for them, and hedged reports are common.\n\n# Everything else in this notebook is ready to run once that file exists.\n# \"\"\"\n\n\n# # ------------------------------- data ----------------------------------------\n# class StudyWindows(Dataset):\n#     \"\"\"Per-study bag of 2.5D windows sampled from all_vols.npy (memmap).\n#     Each window = 3 physically-consecutive slices -> RGB, resized to `res`, in [0,1]\n#     (matches the SSL input pipeline: ToTensor, no ImageNet norm).\"\"\"\n#     def __init__(self, root, ids, id2row, labels, res, k, train, aug=True, norm=\"none\",\n#                  roi=False, roi_mode=\"tight\", roi_pad=0.06, roi_overlap=0.12,\n#                  roi_sbox=None, roi_cen=None):\n#         self.root = root\n#         self.ids = ids\n#         self.id2row = id2row\n#         self.labels = labels              # dict uid -> np.float32[12]\n#         self.res, self.k, self.train, self.aug = res, k, train, aug\n#         self.norm = norm                  # \"none\"=[0,1] (Raptor SSL); \"imagenet\"=DINOv2 stats\n#         self._mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\n#         self._std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n#         self.vols = None; self.masks = None\n#         # --- ROI anatomical joint-crop config ---\n#         self.roi = bool(roi)\n#         self.roi_mode, self.roi_pad, self.roi_overlap = roi_mode, roi_pad, roi_overlap\n#         self.roi_sbox = roi_sbox          # (N,D,4) int16 per-slice tissue bbox, or None\n#         self.roi_cen = roi_cen            # (N,D,2) int16 per-slice joint centroid, or None\n#         if self.roi and (_roi_square_box is None or roi_sbox is None):\n#             raise RuntimeError(\"roi=True but roi_localize or roi_boxes not available\")\n\n#     def __len__(self): return len(self.ids)\n\n#     def _ensure(self):\n#         if self.vols is None:\n#             self.vols, self.masks = _open_corpus()   # (N,D,H,W) view, (N,D) u8\n\n#     def _centers(self, valid, count):\n#         # valid slice indices; window centers must have both neighbors valid & in-range\n#         lo, hi = int(valid.min()), int(valid.max())\n#         cs = [c for c in range(lo + 1, hi) if c - 1 >= lo and c + 1 <= hi]\n#         if not cs: cs = [max(1, min((lo + hi) // 2, self._D - 2))]\n#         if self.train:\n#             reps = count // len(cs) + 1\n#             pool = (cs * reps)\n#             random.shuffle(pool)\n#             return pool[:count]\n#         # eval: evenly spaced deterministic\n#         idx = np.linspace(0, len(cs) - 1, count).round().astype(int)\n#         return [cs[i] for i in idx]\n\n#     def _resize(self, tri):\n#         \"\"\"tri: (3,h,w) float32 [0,1] -> (3,res,res) float32.\"\"\"\n#         t = torch.from_numpy(np.ascontiguousarray(tri))\n#         if t.shape[-1] != self.res or t.shape[-2] != self.res:\n#             t = F.interpolate(t[None], size=(self.res, self.res), mode=\"bilinear\",\n#                               align_corners=False)[0]\n#         return t.numpy()\n\n#     def __getitem__(self, i):\n#         self._ensure()\n#         uid = self.ids[i]; row = self.id2row[uid]\n#         self._D = self.vols.shape[1]\n#         m = self.masks[row]\n#         valid = np.where(m > 0)[0]\n#         if len(valid) < 3: valid = np.arange(min(3, self._D))\n#         # compartment mode emits 2 crops/center -> sample ceil(k/2) centers to keep #windows==k\n#         compart = self.roi and self.roi_mode == \"compartment\"\n#         n_centers = (self.k + 1) // 2 if compart else self.k\n#         cs = self._centers(valid, n_centers)\n#         vol = self.vols[row]  # (D,H,W) u8  (single study read)\n#         tiles = []            # list of (3,res,res) float32\n#         for c in cs:\n#             c = max(1, min(c, self._D - 2))\n#             tri = np.stack([vol[c - 1], vol[c], vol[c + 1]], 0).astype(np.float32) / 255.0  # (3,H,W)\n#             H, W = tri.shape[-2], tri.shape[-1]\n#             if not self.roi:\n#                 tiles.append(self._resize(tri))\n#                 continue\n#             # one joint box from the CENTER slice, applied to all 3 slices (keeps RGB registered)\n#             sq_mode = \"tight\" if compart else self.roi_mode\n#             sq = _roi_square_box(tuple(int(v) for v in self.roi_sbox[row, c]), W=W, H=H,\n#                                  pad=self.roi_pad, mode=sq_mode,\n#                                  centroid=tuple(float(v) for v in self.roi_cen[row, c]))\n#             if compart:\n#                 left, right = _roi_compartments(sq, overlap=self.roi_overlap)\n#                 for box in (left, right):\n#                     x0, y0, x1, y1 = box\n#                     tiles.append(self._resize(tri[:, y0:y1, x0:x1]))\n#             else:\n#                 x0, y0, x1, y1 = sq\n#                 tiles.append(self._resize(tri[:, y0:y1, x0:x1]))\n#         if len(tiles) > self.k:\n#             tiles = tiles[:self.k]\n#         wins = np.stack(tiles, 0)          # (K,3,res,res)\n#         x = torch.from_numpy(wins)\n#         if self.train and self.aug:\n#             # light medical-safe aug: NO flips (laterality is signal); mild intensity jitter\n#             g = 1.0 + (random.random() - 0.5) * 0.20\n#             x = (x * g).clamp(0, 1)\n#         if self.norm == \"imagenet\":        # each backbone at its correct input distribution\n#             x = (x - self._mean) / self._std\n#         y = torch.from_numpy(self.labels[uid])\n#         return x, y\n\n\n# def collate(batch):\n#     xs = torch.stack([b[0] for b in batch])   # (B,K,3,res,res)\n#     ys = torch.stack([b[1] for b in batch])   # (B,12)\n#     return xs, ys\n\n\n# # ------------------------------- model ---------------------------------------\n# def build_backbone(arch=\"vit_small_patch16_224\", pretrained=False):\n#     hybrid = arch.startswith((\"maxvit\", \"maxxvit\", \"coatnet\", \"coat_\", \"convnext\"))\n#     is_vit = (not hybrid) and any(k in arch for k in (\"vit\", \"deit\", \"dinov2\", \"eva\", \"beit\"))\n#     kw = dict(pretrained=pretrained, num_classes=0, in_chans=3)\n#     if is_vit:\n#         kw.update(global_pool=\"token\", dynamic_img_size=True)\n#     else:\n#         kw.update(global_pool=\"avg\")\n#     return timm.create_model(arch, **kw)\n\n\n# def load_raptor(bb, ckpt_path):\n#     if ckpt_path in (\"timm\", \"pretrained\"):\n#         return \"timm-pretrained\"\n#     if ckpt_path in (\"\", \"none\", \"None\"):\n#         print(\"[raptor] RANDOM-INIT control (no SSL weights)\", flush=True)\n#         return \"random-init\"\n#     ck = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n#     st = ck[\"student\"] if \"student\" in ck else ck\n#     bbst = {k[len(\"backbone.\"):]: v for k, v in st.items() if k.startswith(\"backbone.\")}\n#     missing, unexpected = bb.load_state_dict(bbst, strict=False)\n#     ep = ck.get(\"epoch\", \"?\")\n#     print(f\"[raptor] loaded backbone from {os.path.basename(ckpt_path)} (ssl epoch {ep}) | \"\n#           f\"loaded {len(bbst)} tensors, missing {len(missing)}, unexpected {len(unexpected)}\", flush=True)\n#     return f\"{os.path.basename(ckpt_path)}@ep{ep}\"\n\n# class RaptorClassifier(nn.Module):\n#     \"\"\"Raptor encoder + cross-window self-attention + per-diagnosis attention-MIL head.\"\"\"\n#     def __init__(self, backbone, F_dim=384, n=12, drop=0.2, n_heads=4):\n#         super().__init__()\n#         self.backbone = backbone\n#         # --- NEW: single cross-window self-attention layer (windows condition on each other) ---\n#         self.mix_norm = nn.LayerNorm(F_dim)\n#         self.mix = nn.MultiheadAttention(F_dim, n_heads, dropout=drop, batch_first=True)\n#         # --- existing attention-MIL head, unchanged ---\n#         self.norm = nn.LayerNorm(F_dim)\n#         self.att = nn.Sequential(nn.Linear(F_dim, 256), nn.Tanh(), nn.Dropout(drop),\n#                                  nn.Linear(256, n))\n#         self.clsW = nn.Parameter(torch.zeros(n, F_dim))\n#         self.clsb = nn.Parameter(torch.zeros(n))\n#         nn.init.trunc_normal_(self.clsW, std=0.02)\n#         self.n = n\n\n#     def encode(self, x):\n#         B, K = x.shape[:2]\n#         f = self.backbone(x.flatten(0, 1))\n#         return f.view(B, K, -1)\n\n#     def mix_windows(self, feats):\n#         # feats: (B,K,F). Self-attention across the K windows, pre-norm + residual.\n#         h = self.mix_norm(feats)\n#         h, _ = self.mix(h, h, h, need_weights=False)\n#         return feats + h    # residual: transformer can learn identity if it's not helping\n\n#     def head(self, feats):\n#         h = self.norm(feats)\n#         a = self.att(h)\n#         a = torch.softmax(a, dim=1)\n#         pooled = torch.einsum(\"bkn,bkf->bnf\", a, h)\n#         logits = (pooled * self.clsW).sum(-1) + self.clsb\n#         return logits\n\n#     def forward(self, x):\n#         feats = self.encode(x)\n#         feats = self.mix_windows(feats)   # <-- the only new step\n#         return self.head(feats)\n\n\n# # ------------------------------- train ---------------------------------------\n# def main():\n  \n#     ap = argparse.ArgumentParser()\n#     ap.add_argument(\"--ckpt\", default=os.path.join(HERE, \"ckpt\", \"raptor_ssl_last.pt\"))\n#     ap.add_argument(\"--arch\", default=\"vit_small_patch16_224\")\n#     ap.add_argument(\"--res\", type=int, default=224)\n#     ap.add_argument(\"--k\", type=int, default=12)\n#     ap.add_argument(\"--k_eval\", type=int, default=24)\n#     ap.add_argument(\"--epochs\", type=int, default=12)\n#     ap.add_argument(\"--bs\", type=int, default=8)\n#     ap.add_argument(\"--bb_lr\", type=float, default=3e-5)\n#     ap.add_argument(\"--head_lr\", type=float, default=1e-3)\n#     ap.add_argument(\"--wd\", type=float, default=0.02)\n#     ap.add_argument(\"--workers\", type=int, default=4)\n#     ap.add_argument(\"--limit\", type=int, default=0)\n#     ap.add_argument(\"--freeze_blocks\", type=int, default=0)\n#     ap.add_argument(\"--norm\", default=\"none\", choices=[\"none\", \"imagenet\"])\n#     ap.add_argument(\"--grad_ckpt\", action=\"store_true\")\n#     ap.add_argument(\"--tag\", default=\"dev\")\n#     ap.add_argument(\"--labels\", default=None,\n#                     help=\"parquet of training labels: StudyInstanceUID + the twelve finding \"\n#                          \"columns, values 0..1. Not provided with this notebook - see the header.\")\n#     ap.add_argument(\"--smoke\", action=\"store_true\")\n#     ap.add_argument(\"--seed\", type=int, default=42)\n#     # ---- CV fold hook ----\n#     ap.add_argument(\"--folds\", type=int, default=0)\n#     ap.add_argument(\"--fold\", type=int, default=-1)\n#     ap.add_argument(\"--fold_file\", default=None)\n#     # ---- anatomical ROI joint-crop (A/B lever). Default OFF -> identical to baseline. ----\n#     ap.add_argument(\"--roi\", action=\"store_true\")\n#     ap.add_argument(\"--roi_mode\", default=\"compartment\", choices=[\"tight\", \"safe\", \"compartment\"])\n#     ap.add_argument(\"--roi_pad\", type=float, default=0.06)\n#     ap.add_argument(\"--roi_overlap\", type=float, default=0.12)\n#     ap.add_argument(\"--roi_boxes\", default=os.path.join(HERE, \"roi_boxes.npz\"))\n#     a = ap.parse_args()\n#     random.seed(a.seed); np.random.seed(a.seed); torch.manual_seed(a.seed)\n\n#     dev = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n#     scaler = torch.cuda.amp.GradScaler(enabled=dev == \"cuda\")\n\n#     if a.smoke:\n#         a.epochs, a.bs, a.k, a.k_eval, a.limit, a.workers = 2, 4, 4, 6, 40, 0\n#     print(f\"device {dev} | res {a.res} | k {a.k}/{a.k_eval} | bs {a.bs} | tag {a.tag} | \"\n#           f\"roi={a.roi}({a.roi_mode})\", flush=True)\n  \n\n#     # ---- ids / labels ----\n#     ids = _open_ids()\n#     id2row = {u: i for i, u in enumerate(ids)}\n#     idset = set(ids)\n#     _tcsv = os.path.join(RSNA, \"train.csv\")\n#     if not os.path.exists(_tcsv):\n#         _tcsv = _find(\"train.csv\")\n#     if not _tcsv:\n#         raise SystemExit(\"Could not find train.csv - attach the competition data.\")\n#     tr = pd.read_csv(_tcsv); tr[\"StudyInstanceUID\"] = tr[\"StudyInstanceUID\"].astype(str)\n#     gold_df = tr[tr[LAB].notna().all(axis=1)].copy().set_index(\"StudyInstanceUID\")\n#     gold_ids = [u for u in gold_df.index if u in idset]\n#     _lab = a.labels or _find(\"labels_llm_soft.parquet\")\n#     if not _lab or not os.path.exists(_lab):\n#         # Exit cleanly rather than as a failure: running this notebook as published, with no\n#         # labels attached, is the expected path and should read as an explanation, not a crash.\n#         print(_NO_LABELS.format(cols=\", \".join(LAB)), flush=True)\n#         raise SystemExit(0)\n#     soft = pd.read_parquet(_lab)\n#     soft[\"StudyInstanceUID\"] = soft[\"StudyInstanceUID\"].astype(str); soft = soft.set_index(\"StudyInstanceUID\")\n#     goldset = set(gold_ids)\n#     train_ids = [u for u in ids if u in soft.index and u not in goldset]\n#     # ---- CV fold hook: hold out fold `a.fold`, train on the rest ----\n#     oof_ids = []\n#     if a.folds > 0:\n#         assert 0 <= a.fold < a.folds, f\"--fold must be in [0,{a.folds}) when --folds>0\"\n#         fmap = json.load(open(a.fold_file))[\"folds\"]\n#         held = set(u for u in train_ids if fmap.get(u, -1) == a.fold)\n#         oof_ids = [u for u in train_ids if u in held]\n#         train_ids = [u for u in train_ids if u not in held]\n#         print(f\"[cv] fold {a.fold}/{a.folds}: train {len(train_ids)} | OOF held-out {len(oof_ids)} \"\n#               f\"| fold_file {os.path.basename(a.fold_file)}\", flush=True)\n#     if a.limit: train_ids = train_ids[:a.limit]\n#     labels = {u: soft.loc[u, LAB].values.astype(np.float32) for u in train_ids}\n#     for u in oof_ids: labels[u] = soft.loc[u, LAB].values.astype(np.float32)\n#     for u in gold_ids: labels[u] = gold_df.loc[u, LAB].values.astype(np.float32)\n#     print(f\"train {len(train_ids)} | gold-val {len(gold_ids)}\"\n#           + (f\" | oof {len(oof_ids)}\" if oof_ids else \"\"), flush=True)\n\n#     prev = np.clip(np.stack([labels[u] for u in train_ids]).mean(0), 0.03, 0.7)\n#     pw = torch.tensor(np.clip((1 - prev) / prev, 1, 10), dtype=torch.float32, device=dev)\n\n#     # ---- ROI boxes (only when --roi) ----\n#     roi_sbox = roi_cen = None\n#     if a.roi:\n#         rb = np.load(a.roi_boxes)\n#         rb_ids = rb[\"ids\"].astype(str)\n#         if not np.array_equal(rb_ids, ids):\n#             rmap = {u: i for i, u in enumerate(rb_ids)}\n#             missing = [u for u in ids if u not in rmap]\n#             if missing:\n#                 raise RuntimeError(f\"roi_boxes missing {len(missing)} corpus ids (e.g. {missing[:2]})\")\n#             order = np.array([rmap[u] for u in ids])\n#             roi_sbox = rb[\"sbox\"][order]; roi_cen = rb[\"cen\"][order]\n#         else:\n#             roi_sbox = rb[\"sbox\"]; roi_cen = rb[\"cen\"]\n#         print(f\"[roi] ENABLED mode={a.roi_mode} pad={a.roi_pad} overlap={a.roi_overlap} \"\n#               f\"boxes={os.path.basename(a.roi_boxes)} sbox={roi_sbox.shape}\", flush=True)\n#     _roi_kw = dict(roi=a.roi, roi_mode=a.roi_mode, roi_pad=a.roi_pad, roi_overlap=a.roi_overlap,\n#                    roi_sbox=roi_sbox, roi_cen=roi_cen)\n\n#     tds = StudyWindows(HERE, train_ids, id2row, labels, a.res, a.k, train=True, norm=a.norm, **_roi_kw)\n#     vds = StudyWindows(HERE, gold_ids, id2row, labels, a.res, a.k_eval, train=False, norm=a.norm, **_roi_kw)\n#     tl = DataLoader(tds, batch_size=a.bs, shuffle=True, num_workers=a.workers, drop_last=True,\n#                     collate_fn=collate, pin_memory=True, persistent_workers=a.workers > 0)\n#     vl = DataLoader(vds, batch_size=max(2, a.bs // 2), shuffle=False, num_workers=a.workers,\n#                     collate_fn=collate, persistent_workers=a.workers > 0)\n\n#     use_timm = a.ckpt in (\"timm\", \"pretrained\")\n#     bb = build_backbone(a.arch, pretrained=use_timm)\n#     src = load_raptor(bb, a.ckpt)\n#     if use_timm: print(f\"[raptor] timm-pretrained backbone: {a.arch}\", flush=True)\n#     F_dim = bb.num_features\n#     if a.grad_ckpt:\n#         try:\n#             bb.set_grad_checkpointing(True); print(\"[raptor] gradient checkpointing ON\", flush=True)\n#         except Exception as e:\n#             print(f\"[raptor] grad_ckpt unsupported for {a.arch}: {e}\", flush=True)\n#     model = RaptorClassifier(bb, F_dim=F_dim).to(dev)\n#     print(\"Device:\", dev)\n#     print(\"Model device:\", next(model.parameters()).device)\n#     if a.freeze_blocks > 0:\n#         for nm, p in model.backbone.named_parameters():\n#             for b in range(a.freeze_blocks):\n#                 if nm.startswith(f\"blocks.{b}.\"): p.requires_grad = False\n\n#     head_params = [p for n_, p in model.named_parameters() if not n_.startswith(\"backbone.\") and p.requires_grad]\n#     bb_params = [p for n_, p in model.named_parameters() if n_.startswith(\"backbone.\") and p.requires_grad]\n#     opt = torch.optim.AdamW([{\"params\": bb_params, \"lr\": a.bb_lr},\n#                              {\"params\": head_params, \"lr\": a.head_lr}], weight_decay=a.wd)\n#     steps = max(len(tl) * a.epochs, 1)\n#     sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=[a.bb_lr, a.head_lr],\n#                                                 total_steps=steps, pct_start=0.15)\n#     lossf = nn.BCEWithLogitsLoss(pos_weight=pw)\n\n#     @torch.no_grad()\n#     def evaluate():\n#         model.eval(); P = []; Y = []\n#         for x, y in vl:\n#             with torch.autocast(dev, dtype=torch.float16, enabled=dev == \"cuda\"):\n#                 o = torch.sigmoid(model(x.to(dev)).float())\n#             P.append(o.cpu().numpy()); Y.append(y.numpy())\n#         P = np.concatenate(P); Y = np.concatenate(Y)\n#         aucs = {}\n#         for j, name in enumerate(LAB):\n#             if len(set(Y[:, j].astype(int))) > 1:\n#                 aucs[name] = float(roc_auc_score(Y[:, j], P[:, j]))\n#         return float(np.mean(list(aucs.values()))), aucs, P, Y\n\n#     best = 0.0; best_state = None; best_P = None; t0 = time.time(); hist = []\n#     TOPK = 3; topk = []\n#     for ep in range(a.epochs):\n#         model.train(); tot = 0.0\n#         for x, y in tl:\n#             x, y = x.to(dev, non_blocking=True), y.to(dev, non_blocking=True)\n#             opt.zero_grad()\n#             with torch.autocast(dev, dtype=torch.float16, enabled=dev == \"cuda\"):\n#                 logits = model(x)\n#                 loss = lossf(logits.float(), y)\n#             scaler.scale(loss).backward()\n#             scaler.unscale_(opt)\n#             torch.nn.utils.clip_grad_norm_(model.parameters(), 3.0)\n#             scaler.step(opt)\n#             scaler.update()\n#             sched.step(); tot += loss.item()\n#         au, aucs, P, Y = evaluate()\n#         hist.append({\"ep\": ep, \"loss\": tot / len(tl), \"gold_auc\": au})\n#         if au > best:\n#             best = au; best_P = P\n#             best_state = {\"model\": {k: v.detach().cpu() for k, v in model.state_dict().items()},\n#                           \"gold_auc\": au, \"aucs\": aucs, \"src\": src, \"res\": a.res,\n#                           \"arch\": a.arch, \"lab\": LAB, \"epoch\": ep}\n#         if au >= best and best_state is not None:\n#             _tmp = os.path.join(HERE, f\"raptor_ft_{a.tag}.pt.tmp\")\n#             torch.save(best_state, _tmp)\n#             os.replace(_tmp, os.path.join(HERE, f\"raptor_ft_{a.tag}.pt\"))\n#             np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}.npz\"),\n#                      pred=best_P, truth=Y, ids=np.array(gold_ids))\n#             json.dump({\"tag\": a.tag, \"src\": src, \"best_gold_auc\": best,\n#                        \"aucs\": best_state[\"aucs\"], \"hist\": hist, \"res\": a.res,\n#                        \"epochs\": a.epochs, \"bb_lr\": a.bb_lr, \"head_lr\": a.head_lr,\n#                        \"n_train\": len(train_ids), \"n_gold\": len(gold_ids),\n#                        \"roi\": a.roi, \"roi_mode\": a.roi_mode,\n#                        \"partial\": True, \"epochs_done\": ep + 1},\n#                       open(os.path.join(HERE, f\"raptor_ft_{a.tag}.json\"), \"w\"), indent=1)\n#             print(f\"  [ckpt] best-so-far saved at ep{ep} ({best:.4f})\", flush=True)\n#         if len(topk) < TOPK or au > min(t[\"gold_auc\"] for t in topk):\n#             topk.append({\"model\": {k: v.detach().cpu().clone() for k, v in model.state_dict().items()},\n#                          \"gold_auc\": au, \"aucs\": aucs, \"src\": src, \"res\": a.res,\n#                          \"arch\": a.arch, \"lab\": LAB, \"epoch\": ep, \"P\": P})\n#             topk.sort(key=lambda t: -t[\"gold_auc\"])\n#             del topk[TOPK:]\n#         print(f\"ep{ep} loss {tot/len(tl):.3f} | GOLD macro-AUC {au:.4f} (best {best:.4f}) | {time.time()-t0:.0f}s\",\n#               flush=True)\n#     _, aucs, _, Y = evaluate()\n#     print(f\"\\nDONE {src} | BEST GOLD macro-AUC {best:.4f}\", flush=True)\n#     for k, v in (best_state[\"aucs\"] if best_state else aucs).items():\n#         print(f\"   {k:18s} {v:.3f}\", flush=True)\n\n#     if best_state is not None:\n#         torch.save(best_state, os.path.join(HERE, f\"raptor_ft_{a.tag}.pt\"))\n#         np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}.npz\"),\n#                  pred=best_P, truth=Y, ids=np.array(gold_ids))\n#         json.dump({\"tag\": a.tag, \"src\": src, \"best_gold_auc\": best, \"aucs\": best_state[\"aucs\"],\n#                    \"hist\": hist, \"res\": a.res, \"epochs\": a.epochs, \"bb_lr\": a.bb_lr,\n#                    \"head_lr\": a.head_lr, \"n_train\": len(train_ids), \"n_gold\": len(gold_ids),\n#                    \"roi\": a.roi, \"roi_mode\": a.roi_mode, \"partial\": False, \"epochs_done\": a.epochs},\n#                   open(os.path.join(HERE, f\"raptor_ft_{a.tag}.json\"), \"w\"), indent=1)\n#         print(f\"saved raptor_ft_{a.tag}.pt / raptor_gold_{a.tag}.npz / raptor_ft_{a.tag}.json\", flush=True)\n\n#     # ---- CV OOF: predict the held-out fold at the best (gold-selected) weights ----\n#     if a.folds > 0 and oof_ids and best_state is not None:\n#         model.load_state_dict(best_state[\"model\"]); model.eval()\n#         ods = StudyWindows(HERE, oof_ids, id2row, labels, a.res, a.k_eval, train=False, norm=a.norm, **_roi_kw)\n#         ol = DataLoader(ods, batch_size=max(2, a.bs // 2), shuffle=False, num_workers=a.workers,\n#                         collate_fn=collate, persistent_workers=False)\n#         Po = []\n#         with torch.no_grad():\n#             for x, y in ol:\n#                 with torch.autocast(dev, dtype=torch.float16, enabled=dev == \"cuda\"):\n#                     o = torch.sigmoid(model(x.to(dev)).float())\n#                 Po.append(o.cpu().numpy())\n#         Po = np.concatenate(Po)\n#         Yo = np.stack([labels[u] for u in oof_ids]).astype(np.float32)\n#         np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_fold{a.fold}.npz\"),\n#                  pred=Po, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n#         print(f\"[cv] wrote raptor_oof_{a.tag}_fold{a.fold}.npz ({len(oof_ids)} studies; \"\n#               f\"truth = soft labels)\", flush=True)\n\n#         # --- top-K epoch OOF (new) ---\n#         # The epoch-ensemble used to be judged on gold only; that gate is too small to\n#         # resolve the move. Re-run the held-out fold at each retained epoch instead.\n#         oof_by_ep = {best_state[\"epoch\"]: Po}\n#         for t in topk:\n#             e = int(t[\"epoch\"])\n#             if e in oof_by_ep:\n#                 continue\n#             model.load_state_dict(t[\"model\"]); model.eval()\n#             Pe = []\n#             with torch.no_grad():\n#                 for x, y in ol:\n#                     with torch.autocast(dev, dtype=torch.float16, enabled=dev == \"cuda\"):\n#                         Pe.append(torch.sigmoid(model(x.to(dev)).float()).cpu().numpy())\n#             Pe = np.concatenate(Pe); oof_by_ep[e] = Pe\n#             np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_ep{e}_fold{a.fold}.npz\"),\n#                      pred=Pe, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n#             print(f\"[cv] wrote top-K epoch OOF ep{e}\", flush=True)\n#         if len(oof_by_ep) > 1:\n#             Pens = np.mean(list(oof_by_ep.values()), axis=0)\n#             np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_epens_fold{a.fold}.npz\"),\n#                      pred=Pens, truth=Yo, ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n#             print(f\"[cv] wrote epoch-ensemble OOF over epochs {sorted(oof_by_ep)}\", flush=True)\n#         model.load_state_dict(best_state[\"model\"]); model.eval()\n\n#     # --- weight-averaged checkpoint (new) ---\n#     if len(topk) > 1:\n#         import copy\n#         sds = [t[\"model\"] for t in topk]\n#         avg = {}\n#         for k in sds[0]:\n#             v0 = sds[0][k]\n#             if v0.is_floating_point():\n#                 avg[k] = sum(sd[k].double() for sd in sds).div(len(sds)).to(v0.dtype)\n#             else:\n#                 avg[k] = v0.clone()          # e.g. num_batches_tracked\n#         swa_state = {\"model\": avg, \"gold_auc\": None, \"aucs\": {}, \"src\": src, \"res\": a.res,\n#                      \"arch\": a.arch, \"lab\": LAB, \"epoch\": [int(t[\"epoch\"]) for t in topk],\n#                      \"swa_over\": [int(t[\"epoch\"]) for t in topk]}\n#         model.load_state_dict(avg); model.eval()\n#         au_swa, aucs_swa, P_swa, Y_swa = evaluate()\n#         swa_state[\"gold_auc\"] = au_swa; swa_state[\"aucs\"] = aucs_swa\n#         torch.save(swa_state, os.path.join(HERE, f\"raptor_ft_{a.tag}_swa.pt\"))\n#         np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_swa.npz\"),\n#                  pred=P_swa, truth=Y_swa, ids=np.array(gold_ids))\n#         print(f\"SWA over epochs {[int(t['epoch']) for t in topk]} | gold {au_swa:.4f} \"\n#               f\"(best-epoch {best:.4f}) {'BETTER' if au_swa > best else 'no gain'}\", flush=True)\n#         if a.folds > 0 and oof_ids:\n#             ods2 = StudyWindows(HERE, oof_ids, id2row, labels, a.res, a.k_eval, train=False,\n#                                 norm=a.norm, **_roi_kw)\n#             ol2 = DataLoader(ods2, batch_size=max(2, a.bs // 2), shuffle=False,\n#                              num_workers=a.workers, collate_fn=collate, persistent_workers=False)\n#             Ps = []\n#             with torch.no_grad():\n#                 for x, y in ol2:\n#                     with torch.autocast(dev, dtype=torch.float16, enabled=dev == \"cuda\"):\n#                         Ps.append(torch.sigmoid(model(x.to(dev)).float()).cpu().numpy())\n#             Ps = np.concatenate(Ps)\n#             np.savez(os.path.join(HERE, f\"raptor_oof_{a.tag}_swa_fold{a.fold}.npz\"),\n#                      pred=Ps, truth=np.stack([labels[u] for u in oof_ids]).astype(np.float32),\n#                      ids=np.array(oof_ids), fold=a.fold, nfolds=a.folds)\n#             print(f\"[cv] wrote SWA OOF\", flush=True)\n#         model.load_state_dict(best_state[\"model\"]); model.eval()\n\n#     if len(topk) > 1:\n#         eps = [t[\"epoch\"] for t in topk]\n#         for rank, t in enumerate(topk):\n#             if rank == 0:\n#                 continue\n#             P_t = t.pop(\"P\")\n#             np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_ep{t['epoch']}.npz\"),\n#                      pred=P_t, truth=Y, ids=np.array(gold_ids))\n#             t[\"P\"] = P_t\n#         Pens = np.mean([t[\"P\"] for t in topk], axis=0)\n#         ens_aucs = {}\n#         for j, name in enumerate(LAB):\n#             if len(set(Y[:, j].astype(int))) > 1:\n#                 ens_aucs[name] = float(roc_auc_score(Y[:, j], Pens[:, j]))\n#         ens = float(np.mean(list(ens_aucs.values())))\n#         np.savez(os.path.join(HERE, f\"raptor_gold_{a.tag}_epens.npz\"),\n#                  pred=Pens, truth=Y, ids=np.array(gold_ids))\n#         print(f\"TOPK epochs {eps} | best {best:.4f} | epoch-ensemble {ens:.4f} \"\n#               f\"({'BETTER' if ens > best else 'no gain'})\", flush=True)\n\n\n# if __name__ == \"__main__\":\n#     main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.73077Z","iopub.execute_input":"2026-08-26T12:52:46.73112Z","iopub.status.idle":"2026-08-26T12:52:46.752114Z","shell.execute_reply.started":"2026-08-26T12:52:46.731084Z","shell.execute_reply":"2026-08-26T12:52:46.75146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python train_knee.py --smoke --tag smoketest  --ckpt none --labels /kaggle/working/labels_soft.parquet \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.755532Z","iopub.execute_input":"2026-08-26T12:52:46.755851Z","iopub.status.idle":"2026-08-26T12:52:46.766666Z","shell.execute_reply.started":"2026-08-26T12:52:46.755827Z","shell.execute_reply":"2026-08-26T12:52:46.76595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !python train_knee.py \\\n#     --epochs 1 \\\n#     --limit 10 \\\n#     --k 6 \\\n#     --k_eval 4 \\\n#     --bs 2 \\\n#     --workers 0 \\\n#     --tag smoketest \\\n#     --labels /kaggle/working/labels_soft.parquet \\\n#     --ckpt none","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T12:52:46.767425Z","iopub.execute_input":"2026-08-26T12:52:46.767695Z","iopub.status.idle":"2026-08-26T12:52:46.778054Z","shell.execute_reply.started":"2026-08-26T12:52:46.767674Z","shell.execute_reply":"2026-08-26T12:52:46.777306Z"}},"outputs":[],"execution_count":null}]}