{"cells":[{"cell_type":"markdown","metadata":{},"source":"# RSNA Knee — Multilingual Weak Supervision + 2.5D MRI\n\nThis is a self-contained, internet-off Kaggle baseline for the **RSNA Knee Abnormality Detection** competition.\n\nThe training file contains only a small image-labelled subset, while every training study has its original radiology report. This notebook therefore:\n\n1. extracts conservative weak labels from multilingual reports;\n2. replaces weak labels with the official image-derived labels wherever they exist;\n3. loads three central slices from one fluid-sensitive series in each anatomical plane;\n4. trains a compact 9-channel ResNet18; and\n5. writes a format-checked `submission.csv`.\n\nThe report labels are deliberately only *weak supervision*: the competition hosts have clarified that the image-derived labels are authoritative when reports and images disagree. The validation set uses official labels only.\n\n**Runtime target:** GPU, internet disabled, comfortably below the 9-hour competition limit.  \n**License:** Apache 2.0."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"from pathlib import Path\nimport gc\nimport random\nimport re\nimport unicodedata\nimport warnings\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.models import resnet18\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 2026\nIMAGE_SIZE = 160\nSLICES_PER_PLANE = 3\nPLANES = (\"sagittal\", \"coronal\", \"axial\")\nCHANNELS = len(PLANES) * SLICES_PER_PLANE\nBATCH_SIZE = 48\nEPOCHS = 8\nNUM_WORKERS = 2\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nAMP = DEVICE.type == \"cuda\"\n\ndef seed_everything(seed=SEED):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\nseed_everything()\nprint(\"device:\", DEVICE, \"| channels:\", CHANNELS)"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Locate the mounted competition source without assuming which of Kaggle's\n# two mount layouts is active.\ninput_root = Path(\"/kaggle/input\")\nsample_paths = list(input_root.rglob(\"sample_submission.csv\"))\ndata_candidates = [p.parent for p in sample_paths if (p.parent / \"train.csv\").exists()]\nassert data_candidates, \"Attach the RSNA Knee competition data as a Notebook input.\"\nDATA_DIR = data_candidates[0]\n\ntrain = pd.read_csv(DATA_DIR / \"train.csv\")\ntrain_series = pd.read_csv(DATA_DIR / \"train_series.csv\")\ntest = pd.read_csv(DATA_DIR / \"test.csv\")\ntest_series = pd.read_csv(DATA_DIR / \"test_series.csv\")\nsample = pd.read_csv(DATA_DIR / \"sample_submission.csv\")\n\nID = \"StudyInstanceUID\"\nSERIES_ID = \"SeriesInstanceUID\"\nLABELS = [c for c in sample.columns if c != ID]\nREPORT_COL = next(c for c in train.columns if c.casefold() == \"report\")\n\nassert LABELS == [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\", \"Synovitis\",\n    \"Baker's\", \"Contusion\", \"Fracture\"\n]\nassert set(LABELS).issubset(train.columns)\n\ngold_mask = train[LABELS].notna().all(axis=1)\nprint(\"DATA_DIR:\", DATA_DIR)\nprint(\"train/test studies:\", len(train), len(test))\nprint(\"officially labelled studies:\", int(gold_mask.sum()))\ndisplay(train.loc[gold_mask, LABELS].mean().rename(\"gold prevalence\").to_frame())"},{"cell_type":"markdown","metadata":{},"source":"## 1. Conservative multilingual report labelling\n\nThe lexicon covers common English, Spanish, German, Turkish, Dutch, French, Greek, and Cyrillic terminology. It normalizes Unicode first (important for Greek characters) and checks a bidirectional negation window, including post-term Turkish negation.\n\nA missing mention is treated as a weak negative, not as ground truth. Official labels always override these report-derived targets and receive a higher training weight."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Regex fragments. These are intentionally conservative: precision is more\n# useful than recall when the output becomes supervision for an image model.\nLEXICON = {\n    \"ACL\": [\n        r\"\\bacl\\b\", r\"anterior cruciate\", r\"ligament[oa] cruzad[oa] anterior\", r\"\\blca\\b\",\n        r\"vorder(?:es|en) kreuzband\", r\"ön çapraz bağ\", r\"voorste kruisband\",\n        r\"ligament croisé antérieur\", r\"πρόσθι(?:ος|ου) χιαστ\", r\"передн\\w* крестообраз\"\n    ],\n    \"MCL\": [\n        r\"\\bmcl\\b\", r\"medial collateral\", r\"ligament[oa] colateral medial\", r\"\\blcm\\b\",\n        r\"innenband\", r\"medial kollateral\", r\"iç yan bağ\", r\"mediale collaterale band\",\n        r\"ligament collatéral médial\", r\"έσω πλάγι\", r\"медиальн\\w* коллатераль\"\n    ],\n    \"Medial Meniscus\": [\n        r\"medial menisc\", r\"internal menisc\", r\"menisc[oa] (?:medial|intern[oa])\",\n        r\"innenmenisk\", r\"iç menisk\", r\"mediale menisc\", r\"ménisque (?:médial|interne)\",\n        r\"έσω μηνίσ\", r\"медиальн\\w* мениск\"\n    ],\n    \"Lateral Meniscus\": [\n        r\"lateral menisc\", r\"external menisc\", r\"menisc[oa] (?:lateral|extern[oa])\",\n        r\"außenmenisk\", r\"dış menisk\", r\"laterale menisc\", r\"ménisque (?:latéral|externe)\",\n        r\"έξω μηνίσ\", r\"латеральн\\w* мениск\"\n    ],\n    \"Medial OA\": [\n        r\"medial (?:compartment )?(?:osteoarth|arthro|gonarth)\",\n        r\"(?:osteoart|artrosis|gonart).{0,45}(?:medial|intern[oa])\",\n        r\"(?:medial|intern[oa]).{0,45}(?:osteoart|artrosis|gonart)\",\n        r\"(?:medial|innen).{0,45}(?:arthrose|gonarthrose)\",\n        r\"(?:medial|iç).{0,45}(?:osteoartrit|gonartroz)\",\n        r\"(?:médial|interne).{0,45}(?:arthrose|gonarthrose)\",\n        r\"έσω.{0,45}(?:οστεοαρθρ|γονάρθρ)\", r\"медиальн.{0,45}(?:остеоартр|гонартр)\"\n    ],\n    \"Lateral OA\": [\n        r\"lateral (?:compartment )?(?:osteoarth|arthro|gonarth)\",\n        r\"(?:osteoart|artrosis|gonart).{0,45}(?:lateral|extern[oa])\",\n        r\"(?:lateral|extern[oa]).{0,45}(?:osteoart|artrosis|gonart)\",\n        r\"(?:lateral|außen).{0,45}(?:arthrose|gonarthrose)\",\n        r\"(?:lateral|dış).{0,45}(?:osteoartrit|gonartroz)\",\n        r\"(?:latéral|externe).{0,45}(?:arthrose|gonarthrose)\",\n        r\"έξω.{0,45}(?:οστεοαρθρ|γονάρθρ)\", r\"латеральн.{0,45}(?:остеоартр|гонартр)\"\n    ],\n    \"PF OA\": [\n        r\"patellofemoral.{0,45}(?:osteoarth|arthro|degenerat|chondr)\",\n        r\"femoro-?patel.{0,45}(?:osteoart|artrosis|degener|condro)\",\n        r\"retropatell.{0,45}(?:arthrose|degener|chondr)\",\n        r\"patellofemoral.{0,45}(?:osteoartrit|artroz|kondr)\",\n        r\"fémoro-?patell.{0,45}(?:arthrose|dégénér|chondr)\",\n        r\"επιγονατιδομηρια.{0,45}(?:οστεοαρθρ|χονδρο)\",\n        r\"пателлофеморал.{0,45}(?:остеоартр|хондр)\"\n    ],\n    \"Effusion\": [\n        r\"joint effusion\", r\"knee effusion\", r\"derrame articular\", r\"derrame intraarticular\",\n        r\"gelenkerguss\", r\"knieerguss\", r\"eklem efüzy\", r\"eklem sıvı\", r\"gewrichtseffus\",\n        r\"épanchement (?:articulaire|intra-?articulaire)\", r\"υδραρθρ\", r\"αρθρική συλλογή\",\n        r\"выпот (?:в суставе|коленного)\"\n    ],\n    \"Synovitis\": [\n        r\"synovitis\", r\"sinovitis\", r\"synovial hypertroph\", r\"synoviale hypertroph\",\n        r\"synovit\", r\"sinovit\", r\"synovite\", r\"υμενίτι\", r\"синовит\"\n    ],\n    \"Baker's\": [\n        r\"baker'?s? cyst\", r\"popliteal cyst\", r\"quiste de baker\", r\"quiste poplíteo\",\n        r\"baker-?zyste\", r\"poplitealzyste\", r\"baker kist\", r\"kyste de baker\",\n        r\"kyste poplité\", r\"κύστη baker\", r\"подколенн\\w* кист\", r\"кист\\w* бейкер\"\n    ],\n    \"Contusion\": [\n        r\"bone contusion\", r\"bone bruise\", r\"osseous contusion\", r\"contusión ósea\",\n        r\"contusione? ósea\", r\"knochenkontusion\", r\"bone bruise\", r\"kemik kontüzy\",\n        r\"contusion osseuse\", r\"οστικ(?:ή|ο) θλάσ\", r\"ушиб кост\"\n    ],\n    \"Fracture\": [\n        r\"acute fracture\", r\"fracture line\", r\"fractura aguda\", r\"línea de fractura\",\n        r\"akute fraktur\", r\"frakturlinie\", r\"akut kırık\", r\"fraktür hattı\",\n        r\"fracture aiguë\", r\"trait de fracture\", r\"οξύ κάταγμα\", r\"γραμμή κατάγματος\",\n        r\"остр\\w* перелом\", r\"линия перелома\"\n    ],\n}\n\nNEGATIONS = [\n    r\"\\bno\\b\", r\"\\bnot\\b\", r\"without\", r\"absen\", r\"negative for\", r\"no evidence of\",\n    r\"\\bsin\\b\", r\"ningun\", r\"ningún\", r\"niega\", r\"descarta\",\n    r\"\\bkein\\w*\\b\", r\"\\bohne\\b\", r\"nicht nachweis\", r\"ausgeschlossen\",\n    r\"\\byok\\w*\\b\", r\"izlenmedi\", r\"saptanmad\", r\"görülmedi\", r\"mevcut değil\",\n    r\"\\bgeen\\b\", r\"\\bzonder\\b\", r\"afwezig\",\n    r\"\\bpas de\\b\", r\"absence de\", r\"sans signe\",\n    r\"χωρίς\", r\"δεν (?:υπάρχει|παρατηρείται|διαπιστώνεται)\", r\"απουσία\",\n    r\"\\bнет\\b\", r\"\\bбез\\b\", r\"не выяв\", r\"отсутств\"\n]\nNEG_RE = re.compile(\"|\".join(NEGATIONS), re.IGNORECASE)\nTERM_RE = {label: re.compile(\"|\".join(parts), re.IGNORECASE) for label, parts in LEXICON.items()}\n\ndef normalize_report(value):\n    value = \"\" if pd.isna(value) else str(value)\n    value = unicodedata.normalize(\"NFKC\", value).casefold()\n    return re.sub(r\"\\s+\", \" \", value)\n\ndef weak_label(report, label, radius=85):\n    text = normalize_report(report)\n    matches = list(TERM_RE[label].finditer(text))\n    if not matches:\n        return 0.0\n    affirmed = False\n    for match in matches:\n        window = text[max(0, match.start() - radius): min(len(text), match.end() + radius)]\n        if not NEG_RE.search(window):\n            affirmed = True\n            break\n    return float(affirmed)\n\nweak_targets = np.asarray([\n    [weak_label(report, label) for label in LABELS]\n    for report in tqdm(train[REPORT_COL], desc=\"weak-labelling reports\")\n], dtype=np.float32)\n\ngold_y = train.loc[gold_mask, LABELS].to_numpy(np.float32)\nweak_gold = weak_targets[gold_mask.to_numpy()]\nweak_auc = []\nfor j, label in enumerate(LABELS):\n    weak_auc.append(roc_auc_score(gold_y[:, j], weak_gold[:, j]))\nweak_report = pd.DataFrame({\"label\": LABELS, \"report_auc_on_gold\": weak_auc,\n                            \"weak_positive_rate\": weak_targets.mean(0)})\ndisplay(weak_report.sort_values(\"report_auc_on_gold\", ascending=False))\n\ntargets = weak_targets.copy()\ntargets[gold_mask.to_numpy()] = gold_y\nrow_weights = np.ones(len(train), dtype=np.float32)\nrow_weights[gold_mask.to_numpy()] = 6.0"},{"cell_type":"markdown","metadata":{},"source":"## 2. Three-plane 2.5D representation\n\nFor each study, the code selects one preferred fluid-sensitive series in each anatomical plane and samples three evenly spaced slices. Percentile clipping, DICOM rescale parameters, and `MONOCHROME1` inversion are handled explicitly.\n\nThe arrays are cached as float16 memory maps in `/kaggle/working`, so decoding occurs once and training remains fast."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def select_plane_series(series_frame):\n    s = series_frame.copy()\n    s[\"_plane\"] = s[\"Anatomical_Plane\"].fillna(\"\").str.casefold()\n    fluid = pd.to_numeric(s.get(\"Fluid_Sensitive\", 0), errors=\"coerce\").fillna(0)\n    fat = pd.to_numeric(s.get(\"Fat_Suppression\", 0), errors=\"coerce\").fillna(0)\n    s[\"_rank\"] = fluid * 2 + fat\n    records = []\n    for study_uid, group in s.groupby(ID, sort=False):\n        row = {ID: study_uid}\n        for plane in PLANES:\n            candidates = group[group[\"_plane\"].str.contains(plane, regex=False)]\n            if len(candidates):\n                candidates = candidates.sort_values([\"_rank\", SERIES_ID], ascending=[False, True])\n                row[f\"{plane}_series\"] = candidates.iloc[0][SERIES_ID]\n            else:\n                row[f\"{plane}_series\"] = None\n        records.append(row)\n    return pd.DataFrame(records)\n\ntrain_slots = select_plane_series(train_series)\ntest_slots = select_plane_series(test_series)\ntrain_meta = train[[ID]].merge(train_slots, on=ID, how=\"left\", validate=\"one_to_one\")\ntest_meta = test[[ID]].merge(test_slots, on=ID, how=\"left\", validate=\"one_to_one\")\n\nmissing_train = train_meta[[f\"{p}_series\" for p in PLANES]].isna().sum()\nmissing_test = test_meta[[f\"{p}_series\" for p in PLANES]].isna().sum()\ndisplay(pd.DataFrame({\"train_missing\": missing_train, \"test_missing\": missing_test}))"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def resize_2d(image, size=IMAGE_SIZE):\n    tensor = torch.from_numpy(np.ascontiguousarray(image))[None, None].float()\n    tensor = nn.functional.interpolate(tensor, size=(size, size), mode=\"bilinear\", align_corners=False)\n    return tensor[0, 0].numpy()\n\ndef decode_slice(path):\n    ds = pydicom.dcmread(path)\n    image = ds.pixel_array.astype(np.float32)\n    image = image * float(getattr(ds, \"RescaleSlope\", 1.0)) + float(getattr(ds, \"RescaleIntercept\", 0.0))\n    finite = np.isfinite(image)\n    if not finite.any():\n        return np.zeros((IMAGE_SIZE, IMAGE_SIZE), np.float32)\n    fill = float(np.median(image[finite]))\n    image = np.nan_to_num(image, nan=fill, posinf=fill, neginf=fill)\n    lo, hi = np.percentile(image, (1.0, 99.0))\n    image = np.clip((image - lo) / max(hi - lo, 1e-6), 0.0, 1.0)\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        image = 1.0 - image\n    return resize_2d(image)\n\ndef decode_series(series_dir, count=SLICES_PER_PLANE):\n    files = sorted(Path(series_dir).glob(\"*.dcm\"))\n    if not files:\n        return np.zeros((count, IMAGE_SIZE, IMAGE_SIZE), np.float32)\n    indices = np.linspace(0, len(files) - 1, count + 2)[1:-1].round().astype(int)\n    images = []\n    for idx in indices:\n        try:\n            images.append(decode_slice(files[int(idx)]))\n        except Exception:\n            images.append(np.zeros((IMAGE_SIZE, IMAGE_SIZE), np.float32))\n    return np.stack(images)\n\ndef decode_study(row, image_root):\n    channels = []\n    for plane in PLANES:\n        series_uid = row[f\"{plane}_series\"]\n        if pd.isna(series_uid):\n            channels.append(np.zeros((SLICES_PER_PLANE, IMAGE_SIZE, IMAGE_SIZE), np.float32))\n        else:\n            series_dir = Path(image_root) / str(row[ID]) / str(series_uid)\n            channels.append(decode_series(series_dir))\n    return np.concatenate(channels).astype(np.float16)\n\ndef build_cache(frame, image_root, cache_path):\n    cache_path = Path(cache_path)\n    expected = (len(frame), CHANNELS, IMAGE_SIZE, IMAGE_SIZE)\n    if cache_path.exists():\n        cached = np.load(cache_path, mmap_mode=\"r\")\n        if cached.shape == expected:\n            print(\"reusing\", cache_path, cached.shape)\n            return cached\n        del cached\n        cache_path.unlink()\n    mmap = np.lib.format.open_memmap(cache_path, mode=\"w+\", dtype=np.float16, shape=expected)\n    for i, (_, row) in enumerate(tqdm(frame.iterrows(), total=len(frame), desc=f\"decoding {cache_path.name}\")):\n        mmap[i] = decode_study(row, image_root)\n    mmap.flush()\n    del mmap\n    gc.collect()\n    return np.load(cache_path, mmap_mode=\"r\")\n\ntrain_cache = build_cache(train_meta, DATA_DIR / \"train_series\", \"/kaggle/working/rsna_train_2p5d.npy\")\ntest_cache = build_cache(test_meta, DATA_DIR / \"test_series\", \"/kaggle/working/rsna_test_2p5d.npy\")\nprint(\"train cache:\", train_cache.shape, train_cache.dtype)"},{"cell_type":"markdown","metadata":{},"source":"## 3. Gold-only validation and weighted training\n\nA held-out portion of the officially labelled studies is removed entirely from training. All remaining studies are used, with official rows weighted more heavily than report-derived rows. Validation macro AUC ignores a class only if the small split happens to contain a single target value for it."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"gold_indices = np.flatnonzero(gold_mask.to_numpy())\n# Stratifying the very small gold subset across 12 labels is unstable; use a\n# fixed split and transparently skip any one-class validation target.\ngold_train_idx, gold_valid_idx = train_test_split(\n    gold_indices, test_size=0.25, random_state=SEED, shuffle=True\n)\nvalid_set = set(gold_valid_idx.tolist())\ntrain_indices = np.asarray([i for i in range(len(train)) if i not in valid_set], dtype=np.int64)\n\nclass_positive = targets[train_indices].sum(0)\npos_weight = (len(train_indices) - class_positive) / np.maximum(class_positive, 1.0)\npos_weight = torch.tensor(np.clip(pos_weight, 0.5, 12.0), dtype=torch.float32, device=DEVICE)\n\nclass KneeArrayDataset(Dataset):\n    def __init__(self, array, indices, y=None, weights=None, augment=False):\n        self.array = array\n        self.indices = np.asarray(indices)\n        self.y = y\n        self.weights = weights\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, item):\n        idx = int(self.indices[item])\n        x = torch.from_numpy(np.array(self.array[idx], dtype=np.float32, copy=True))\n        if self.augment:\n            if torch.rand(()) < 0.5:\n                x = torch.flip(x, dims=(-1,))\n            gain = 0.90 + 0.20 * torch.rand(())\n            bias = -0.04 + 0.08 * torch.rand(())\n            x = (x * gain + bias).clamp(0, 1)\n        if self.y is None:\n            return x, idx\n        y = torch.from_numpy(self.y[idx].astype(np.float32))\n        weight = torch.tensor(float(self.weights[idx]), dtype=torch.float32)\n        return x, y, weight\n\ntrain_loader = DataLoader(\n    KneeArrayDataset(train_cache, train_indices, targets, row_weights, augment=True),\n    batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True,\n    persistent_workers=NUM_WORKERS > 0\n)\nvalid_loader = DataLoader(\n    KneeArrayDataset(train_cache, gold_valid_idx, targets, row_weights, augment=False),\n    batch_size=BATCH_SIZE * 2, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True\n)\n\nprint(\"training rows:\", len(train_indices), \"| gold validation rows:\", len(gold_valid_idx))\nprint(\"pos_weight:\", dict(zip(LABELS, pos_weight.detach().cpu().numpy().round(2))))"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"class KneeResNet(nn.Module):\n    def __init__(self, in_channels, outputs):\n        super().__init__()\n        self.net = resnet18(weights=None)\n        self.net.conv1 = nn.Conv2d(in_channels, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        features = self.net.fc.in_features\n        self.net.fc = nn.Sequential(nn.Dropout(0.25), nn.Linear(features, outputs))\n\n    def forward(self, x):\n        return self.net(x)\n\nmodel = KneeResNet(CHANNELS, len(LABELS)).to(DEVICE)\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight, reduction=\"none\")\noptimizer = torch.optim.AdamW(model.parameters(), lr=8e-4, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=5e-6)\nscaler = torch.cuda.amp.GradScaler(enabled=AMP)\n\ndef validation_predictions(loader):\n    model.eval()\n    predictions, truths = [], []\n    with torch.no_grad():\n        for x, y, _ in loader:\n            x = x.to(DEVICE, non_blocking=True)\n            with torch.cuda.amp.autocast(enabled=AMP):\n                logits = model(x)\n            predictions.append(torch.sigmoid(logits).cpu().numpy())\n            truths.append(y.numpy())\n    return np.concatenate(truths), np.concatenate(predictions)\n\ndef macro_auc(y_true, y_pred):\n    scores = []\n    for j in range(y_true.shape[1]):\n        scores.append(roc_auc_score(y_true[:, j], y_pred[:, j])\n                      if np.unique(y_true[:, j]).size == 2 else np.nan)\n    return float(np.nanmean(scores)), np.asarray(scores)\n\nbest_score = -np.inf\nbest_state = None\nhistory = []\n\nfor epoch in range(1, EPOCHS + 1):\n    model.train()\n    running = []\n    for x, y, weight in train_loader:\n        x = x.to(DEVICE, non_blocking=True)\n        y = y.to(DEVICE, non_blocking=True)\n        weight = weight.to(DEVICE, non_blocking=True)\n        optimizer.zero_grad(set_to_none=True)\n        with torch.cuda.amp.autocast(enabled=AMP):\n            logits = model(x)\n            per_row = criterion(logits, y).mean(1)\n            loss = (per_row * weight).sum() / weight.sum()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        running.append(float(loss.detach().cpu()))\n\n    y_valid, p_valid = validation_predictions(valid_loader)\n    score, per_class = macro_auc(y_valid, p_valid)\n    scheduler.step()\n    row = {\"epoch\": epoch, \"loss\": float(np.mean(running)), \"gold_macro_auc\": score}\n    history.append(row)\n    print(row)\n    if score > best_score:\n        best_score = score\n        best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\nmodel.load_state_dict(best_state)\nhistory = pd.DataFrame(history)\ndisplay(history)\ndisplay(pd.DataFrame({\"label\": LABELS, \"gold_valid_auc\": per_class}).sort_values(\"gold_valid_auc\", ascending=False))\nprint(\"best gold macro AUC:\", round(best_score, 4))"},{"cell_type":"markdown","metadata":{},"source":"## 4. Test-time augmentation and submission\n\nPredictions are averaged between the original and horizontally flipped images. The final cell verifies row order, column order, finiteness, range, and exact filename before finishing."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"test_indices = np.arange(len(test), dtype=np.int64)\ntest_loader = DataLoader(\n    KneeArrayDataset(test_cache, test_indices),\n    batch_size=BATCH_SIZE * 2, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True\n)\n\nmodel.eval()\ntest_predictions = []\nwith torch.no_grad():\n    for x, _ in tqdm(test_loader, desc=\"inference\"):\n        x = x.to(DEVICE, non_blocking=True)\n        with torch.cuda.amp.autocast(enabled=AMP):\n            logits_a = model(x)\n            logits_b = model(torch.flip(x, dims=(-1,)))\n            probs = 0.5 * (torch.sigmoid(logits_a) + torch.sigmoid(logits_b))\n        test_predictions.append(probs.cpu().numpy())\n\ntest_predictions = np.concatenate(test_predictions)\nsubmission = pd.DataFrame(test_predictions, columns=LABELS)\nsubmission.insert(0, ID, test[ID].to_numpy())\n\nassert submission.columns.tolist() == sample.columns.tolist()\nassert submission[ID].tolist() == sample[ID].tolist(), \"Test/sample row order mismatch.\"\nassert submission.shape == sample.shape\nassert np.isfinite(submission[LABELS].to_numpy()).all()\nassert submission[LABELS].to_numpy().min() >= 0\nassert submission[LABELS].to_numpy().max() <= 1\n\nsubmission.to_csv(\"submission.csv\", index=False)\ndisplay(submission.head())\nprint(\"Wrote submission.csv:\", submission.shape)\nprint(\"mean predictions:\")\ndisplay(submission[LABELS].mean().rename(\"mean_probability\").to_frame())"},{"cell_type":"markdown","metadata":{},"source":"### Notes for iteration\n\n- The most valuable next improvement is stronger report supervision or externally published model weights, added as a Kaggle Dataset/Model input while keeping internet disabled.\n- Keep the official image labels authoritative; report labels are noisy by design.\n- For a serious leaderboard run, cross-validation across the 58 official rows is more informative than trusting a single small holdout.\n- This notebook intentionally avoids private data, network calls, and hidden dependencies."}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"kaggle":{"accelerator":"gpu","dataSources":[]}},"nbformat":4,"nbformat_minor":5}