{"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":"import os\nimport gc\nimport time\nimport random\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as tvm\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\n\ntry:\n    from tqdm.auto import tqdm\nexcept ImportError:\n    tqdm = lambda x, **kwargs: x\n\nRUN_START = time.time()\n\ndef elapsed():\n    h = (time.time() - RUN_START) / 3600\n    return f\"{h:.2f}h\"\n\n# ================================================================\n# 0. ROBUST PATH FINDING\n# ================================================================\n\ndef find_file(base, keywords, exts=('.pth', '.pt')):\n    if not os.path.isdir(base):\n        return None\n    for root, _, files in os.walk(base):\n        for f in files:\n            if f.endswith(exts) and any(k.lower() in f.lower() for k in keywords):\n                return os.path.join(root, f)\n    return None\n\nBASE = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\n# Proven V9 paths\nEFFICIENTNET_WEIGHTS_PATH = (\n    \"/kaggle/input/models/elizavetanew/efficientnet_b0/pytorch/\"\n    \"efficientnet_b0_rwightman-7f5810bc/1/efficientnet_b0_rwightman-7f5810bc.pth\"\n)\nRADIMAGENET_WEIGHTS_PATH = (\n    \"/kaggle/input/datasets/marwanmath/resnet-50-radimagenet-marwan/ResNet50.pt\"\n)\nLLM_LABELS_PATH = (\n    \"/kaggle/input/datasets/stevenleehans/rsna-knee-llm-report-labels/llm_labels_v2.csv\"\n)\n\n# Auto-fallbacks\nif not os.path.exists(EFFICIENTNET_WEIGHTS_PATH):\n    EFFICIENTNET_WEIGHTS_PATH = find_file(\"/kaggle/input\", [\"efficientnet_b0\", \"efficientnet-b0\"])\nif not os.path.exists(RADIMAGENET_WEIGHTS_PATH):\n    RADIMAGENET_WEIGHTS_PATH = find_file(\"/kaggle/input\", [\"resnet50\", \"radimagenet\"])\nif not os.path.exists(LLM_LABELS_PATH):\n    LLM_LABELS_PATH = find_file(\"/kaggle/input\", [\"llm\", \"knee\"], exts=('.csv',))\n\n# NEW: ConvNeXt Base\nCONVEXT_PATH = find_file(\"/kaggle/input\", [\"convnext\"], exts=('.pth', '.pt'))\n\nprint(f\"EfficientNet B0: {EFFICIENTNET_WEIGHTS_PATH}\")\nprint(f\"RadImageNet:     {RADIMAGENET_WEIGHTS_PATH}\")\nprint(f\"LLM Labels:      {LLM_LABELS_PATH}\")\nprint(f\"ConvNeXt Base:   {CONVEXT_PATH}\")\n\n# ================================================================\n# 1. CONFIG (3 BACKBONES, FITS IN 9H)\n# ================================================================\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nLABEL_COLS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\nN_LABELS = len(LABEL_COLS)\n\nSLICES_PER_STUDY = 12\nIMG_SIZE = 224\nN_FOLDS = 5\nBATCH_SIZE = 8\nNUM_WORKERS = 2\n\nPSEUDO_LABEL_WEIGHT = 0.35\nGOLD_LABEL_WEIGHT = 2.0\n\n# Epochs tuned to fit 3 backbones in ~8.5 hours\nBACKBONE_CONFIG = {\n    \"efficientnet_b0\":      {\"epochs\": 5, \"lr\": 3e-4},\n    \"resnet50_radimagenet\": {\"epochs\": 6, \"lr\": 1.5e-4},\n}\n\nif CONVEXT_PATH:\n    BACKBONE_CONFIG[\"convnext_base\"] = {\"epochs\": 4, \"lr\": 1e-4}\n\nest_hours = sum(v[\"epochs\"] for v in BACKBONE_CONFIG.values()) * N_FOLDS * 0.075\nprint(\"\\n\" + \"=\" * 70)\nprint(\"RSNA KNEE -- V11: 3-BACKBONE ENSEMBLE\")\nprint(\"=\" * 70)\nprint(f\"Backbones ({len(BACKBONE_CONFIG)}): {list(BACKBONE_CONFIG.keys())}\")\nprint(f\"Est. training time: ~{est_hours:.1f}h | {elapsed()}\")\nprint(\"=\" * 70)\n\nCACHE_ROOT = \"/kaggle/working/rsna_knee_cache_v11\"\nTRAIN_CACHE = os.path.join(CACHE_ROOT, f\"train_{SLICES_PER_STUDY}_{IMG_SIZE}\")\nTEST_CACHE = os.path.join(CACHE_ROOT, f\"test_{SLICES_PER_STUDY}_{IMG_SIZE}\")\nos.makedirs(TRAIN_CACHE, exist_ok=True)\nos.makedirs(TEST_CACHE, exist_ok=True)\n\n# ================================================================\n# 2. LOAD DATA + LLM PSEUDO-LABELS (EXACTLY V9)\n# ================================================================\n\ntrain = pd.read_csv(f\"{BASE}/train.csv\")\ntrain_series = pd.read_csv(f\"{BASE}/train_series.csv\")\ntest = pd.read_csv(f\"{BASE}/test.csv\")\ntest_series = pd.read_csv(f\"{BASE}/test_series.csv\")\nsample_sub = pd.read_csv(f\"{BASE}/sample_submission.csv\")\n\ngold_mask = train[LABEL_COLS[0]].notna()\n\nif LLM_LABELS_PATH and os.path.exists(LLM_LABELS_PATH):\n    llm_labels = pd.read_csv(LLM_LABELS_PATH)\n    llm_labels = llm_labels.rename(columns={c: f\"{c}_llm\" for c in LABEL_COLS})\n    train = train.merge(llm_labels, on=\"StudyInstanceUID\", how=\"left\")\n    print(f\"\\nLLM labels loaded: {len(llm_labels)} studies\")\nelse:\n    raise RuntimeError(\"LLM labels NOT FOUND. V11 requires them.\")\n\nfinal_labels = train[[\"StudyInstanceUID\"]].copy()\nfor col in LABEL_COLS:\n    gold_col = train[col]\n    llm_col = train[f\"{col}_llm\"]\n    final_labels[col] = gold_col.where(gold_col.notna(), llm_col)\n\nfor col in LABEL_COLS:\n    llm_confidence = (2 * (train[f\"{col}_llm\"] - 0.5).abs()).clip(0, 1)\n    weight_col = pd.Series(\n        np.where(gold_mask, GOLD_LABEL_WEIGHT, PSEUDO_LABEL_WEIGHT * llm_confidence),\n        index=train.index, dtype=np.float32\n    )\n    final_labels[f\"{col}_weight\"] = weight_col.values\n\nfinal_labels[\"is_gold\"] = gold_mask.values\nhas_any_label = final_labels[LABEL_COLS].notna().any(axis=1)\nfinal_labels = final_labels[has_any_label].reset_index(drop=True)\n\ngold_df = final_labels[final_labels[\"is_gold\"]].reset_index(drop=True)\npseudo_df = final_labels[~final_labels[\"is_gold\"]].reset_index(drop=True)\nprint(f\"\\nGold studies: {len(gold_df)} | Pseudo studies: {len(pseudo_df)}\")\n\nstrat_col = gold_df[\"Fracture\"].fillna(0).astype(int)\nskf = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nsplits = list(skf.split(gold_df, strat_col))\n\nlabel_block_binary = (final_labels[LABEL_COLS].astype(float) > 0.5).astype(float)\npos_counts = label_block_binary.sum()\nneg_counts = (1 - label_block_binary).sum()\npos_weight_values = (neg_counts / pos_counts.clip(lower=1)).values\npos_weight = torch.tensor(pos_weight_values, dtype=torch.float32).to(DEVICE)\n\n# ================================================================\n# 3. SERIES LOOKUP + CACHING (EXACTLY V9)\n# ================================================================\n\ndef build_series_lookup(series_df):\n    lookup = {}\n    for row in series_df.itertuples(index=False):\n        uid = str(row.StudyInstanceUID)\n        plane = getattr(row, 'Anatomical_Plane', 'Unknown')\n        lookup.setdefault(uid, []).append((str(row.SeriesInstanceUID), plane))\n    return lookup\n\nTRAIN_SERIES_LOOKUP = build_series_lookup(train_series)\nTEST_SERIES_LOOKUP = build_series_lookup(test_series)\n\ndef get_study_slice_paths(study_uid, series_lookup, split, n_slices):\n    study_uid = str(study_uid)\n    study_dir = os.path.join(BASE, split, study_uid)\n    series_entries = series_lookup.get(study_uid, [])\n    by_plane = {}\n    for series_uid, plane in series_entries:\n        series_dir = os.path.join(study_dir, series_uid)\n        if not os.path.isdir(series_dir):\n            continue\n        try:\n            files = sorted(f for f in os.listdir(series_dir) if not f.startswith(\".\"))\n        except Exception:\n            continue\n        if not files:\n            continue\n        by_plane.setdefault(plane, []).append((series_dir, files))\n    if not by_plane:\n        return []\n    planes = list(by_plane.keys())\n    n_planes = len(planes)\n    base_quota = n_slices // n_planes\n    remainder = n_slices % n_planes\n    all_slices = []\n    for i, plane in enumerate(planes):\n        quota = base_quota + (1 if i < remainder else 0)\n        if quota == 0:\n            continue\n        plane_candidates = []\n        for series_dir, files in by_plane[plane]:\n            take = min(4, len(files))\n            idxs = np.linspace(0, len(files) - 1, take).astype(int)\n            for idx in idxs:\n                plane_candidates.append(os.path.join(series_dir, files[idx]))\n        if not plane_candidates:\n            continue\n        if len(plane_candidates) >= quota:\n            pick_idxs = np.linspace(0, len(plane_candidates) - 1, quota).astype(int)\n            all_slices.extend(plane_candidates[i] for i in pick_idxs)\n        else:\n            all_slices.extend(plane_candidates)\n    if not all_slices:\n        return []\n    if len(all_slices) >= n_slices:\n        pick_idxs = np.linspace(0, len(all_slices) - 1, n_slices).astype(int)\n        chosen = [all_slices[i] for i in pick_idxs]\n    else:\n        chosen = all_slices + [all_slices[-1]] * (n_slices - len(all_slices))\n    return chosen\n\ndef process_dicom(path):\n    try:\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        lo, hi = np.percentile(img, [1, 99])\n        img = np.clip((img - lo) / (hi - lo + 1e-6), 0, 1)\n        return img.astype(np.float32)\n    except Exception:\n        return np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n\ndef cache_one_study(study_uid, series_lookup, split, cache_dir):\n    study_uid = str(study_uid)\n    cache_path = os.path.join(cache_dir, f\"{study_uid}.npy\")\n    if os.path.exists(cache_path):\n        return cache_path\n    paths = get_study_slice_paths(study_uid, series_lookup, split, SLICES_PER_STUDY)\n    out = np.zeros((SLICES_PER_STUDY, IMG_SIZE, IMG_SIZE), dtype=np.float32)\n    if paths:\n        for i, path in enumerate(paths):\n            out[i] = process_dicom(path)\n    np.save(cache_path, out)\n    return cache_path\n\ndef build_cache(study_uids, series_lookup, split, cache_dir, name):\n    print(f\"\\nBuilding {name} cache... [{elapsed()}]\")\n    missing = 0\n    for uid in tqdm(study_uids, desc=f\"Caching {name}\"):\n        cache_path = os.path.join(cache_dir, f\"{str(uid)}.npy\")\n        if not os.path.exists(cache_path):\n            missing += 1\n            cache_one_study(uid, series_lookup, split, cache_dir)\n    print(f\"{name} cache complete. Newly processed: {missing} [{elapsed()}]\")\n\ntrain_uids = final_labels[\"StudyInstanceUID\"].astype(str).tolist()\ntest_uids = test[\"StudyInstanceUID\"].astype(str).tolist()\n\nbuild_cache(train_uids, TRAIN_SERIES_LOOKUP, \"train_series\", TRAIN_CACHE, \"train\")\nbuild_cache(test_uids, TEST_SERIES_LOOKUP, \"test_series\", TEST_CACHE, \"test\")\n\n# ================================================================\n# 4. DATASET (EXACTLY V9)\n# ================================================================\n\nclass KneeStudyDataset(Dataset):\n    def __init__(self, df, cache_dir, training=True):\n        self.df = df.reset_index(drop=True).copy()\n        self.cache_dir = cache_dir\n        self.training = training\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        uid = str(row[\"StudyInstanceUID\"])\n        slices = np.load(os.path.join(self.cache_dir, f\"{uid}.npy\")).astype(np.float32)\n\n        if not self.training:\n            return torch.from_numpy(slices).unsqueeze(1), uid\n\n        label_values = row[LABEL_COLS].to_numpy(dtype=object)\n        labels = np.zeros(N_LABELS, dtype=np.float32)\n        label_mask = np.zeros(N_LABELS, dtype=np.float32)\n        for i, value in enumerate(label_values):\n            if pd.notna(value):\n                labels[i] = float(value)\n                label_mask[i] = 1.0\n        weight = row[[f\"{col}_weight\" for col in LABEL_COLS]].to_numpy(dtype=np.float32)\n\n        return (\n            torch.from_numpy(slices).unsqueeze(1),\n            torch.from_numpy(labels),\n            torch.from_numpy(label_mask),\n            torch.from_numpy(weight),\n        )\n\n# ================================================================\n# 5. MODEL (V9 + ConvNeXt support)\n# ================================================================\n\nclass AttentionMILModel(nn.Module):\n    def __init__(self, n_labels=N_LABELS, backbone_choice=\"efficientnet_b0\", pretrained_path=None):\n        super().__init__()\n        self.backbone_choice = backbone_choice\n\n        if backbone_choice == \"resnet50_radimagenet\":\n            backbone = tvm.resnet50(weights=None)\n            if pretrained_path and os.path.exists(pretrained_path):\n                checkpoint = torch.load(pretrained_path, map_location=\"cpu\")\n                prefix_map = {\"0\": \"conv1\", \"1\": \"bn1\", \"4\": \"layer1\",\n                              \"5\": \"layer2\", \"6\": \"layer3\", \"7\": \"layer4\"}\n                remapped = {}\n                for key, value in checkpoint.items():\n                    parts = key.split(\".\")\n                    idx = parts[1]\n                    if idx not in prefix_map:\n                        continue\n                    new_key = (prefix_map[idx] + \".\" + \".\".join(parts[2:])\n                               if len(parts) > 2 else prefix_map[idx])\n                    remapped[new_key] = value\n                backbone.load_state_dict(remapped, strict=False)\n            self.encoder = nn.Sequential(*list(backbone.children())[:-2])\n            self.encoder_type = \"cnn\"\n            feat_dim = 2048\n\n        elif backbone_choice == \"convnext_base\":\n            try:\n                import timm\n                backbone = timm.create_model('convnext_base.fb_in22k', pretrained=False, num_classes=0, in_chans=3)\n                if pretrained_path and os.path.exists(pretrained_path):\n                    state = torch.load(pretrained_path, map_location=\"cpu\")\n                    backbone.load_state_dict(state, strict=False)\n                self.encoder = backbone\n                self.encoder_type = \"timm\"\n                feat_dim = backbone.num_features\n            except Exception as e:\n                print(f\"ConvNeXt load failed: {e}. Falling back to EfficientNet-B0.\")\n                backbone = tvm.efficientnet_b0(weights=None)\n                self.encoder = backbone.features\n                self.encoder_type = \"cnn\"\n                feat_dim = 1280\n\n        else:  # efficientnet_b0\n            backbone = tvm.efficientnet_b0(weights=None)\n            if pretrained_path and os.path.exists(pretrained_path):\n                state = torch.load(pretrained_path, map_location=\"cpu\")\n                if isinstance(state, dict) and \"state_dict\" in state:\n                    state = state[\"state_dict\"]\n                cleaned = {k[7:] if k.startswith(\"module.\") else k: v for k, v in state.items()}\n                backbone.load_state_dict(cleaned, strict=False)\n            self.encoder = backbone.features\n            self.encoder_type = \"cnn\"\n            feat_dim = 1280\n\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.attention = nn.Sequential(nn.Linear(feat_dim, 128), nn.Tanh(), nn.Linear(128, 1))\n        self.head = nn.Sequential(\n            nn.Linear(feat_dim, 256), nn.ReLU(), nn.Dropout(0.3), nn.Linear(256, n_labels)\n        )\n\n    def forward(self, x):\n        B, S, C, H, W = x.shape\n        x = x.reshape(B * S, C, H, W).repeat(1, 3, 1, 1)\n        feats = self.encoder(x)\n\n        if self.encoder_type == \"cnn\":\n            feats = self.pool(feats).flatten(1)\n        # timm models (ConvNeXt) already return pooled (B*S, feat_dim)\n\n        feats = feats.reshape(B, S, -1)\n        attn_weights = torch.softmax(self.attention(feats), dim=1)\n        study_feat = (feats * attn_weights).sum(dim=1)\n        return self.head(study_feat)\n\ndef masked_weighted_bce(logits, labels, label_mask, label_weight, pos_weight):\n    loss = F.binary_cross_entropy_with_logits(logits, labels, pos_weight=pos_weight, reduction=\"none\")\n    loss = loss * label_mask * label_weight\n    denom = (label_mask * label_weight).sum(dim=1).clamp(min=1e-6)\n    return (loss.sum(dim=1) / denom).mean()\n\nUSE_AMP = DEVICE.type == \"cuda\"\n\ndef autocast_context():\n    if USE_AMP:\n        return torch.amp.autocast(\"cuda\")\n    return torch.autocast(device_type=\"cpu\", enabled=False)\n\ndef make_loader(dataset, shuffle):\n    kwargs = {\"batch_size\": BATCH_SIZE, \"shuffle\": shuffle, \"num_workers\": NUM_WORKERS,\n              \"pin_memory\": (DEVICE.type == \"cuda\")}\n    if NUM_WORKERS > 0:\n        kwargs[\"persistent_workers\"] = True\n    return DataLoader(dataset, **kwargs)\n\ndef compute_val_macro_auc(fold_val, val_preds):\n    aucs = []\n    per_label = {}\n    for i, col in enumerate(LABEL_COLS):\n        values = fold_val[col].astype(float).values\n        valid = ~np.isnan(values)\n        if valid.sum() > 1 and len(np.unique(values[valid])) > 1:\n            auc = roc_auc_score(values[valid], val_preds[valid, i])\n            aucs.append(auc)\n            per_label[col] = auc\n    macro = np.mean(aucs) if aucs else float(\"nan\")\n    return macro, per_label\n\ndef checkpoint_paths(backbone, fold):\n    if backbone == \"efficientnet_b0\":\n        return (f\"/kaggle/working/v11_model_fold{fold}.pt\",\n                f\"/kaggle/working/v11_meta_fold{fold}.npz\")\n    return (f\"/kaggle/working/v11_model_fold{fold}_{backbone}.pt\",\n            f\"/kaggle/working/v11_meta_fold{fold}_{backbone}.npz\")\n\ndef pretrained_path_for(backbone):\n    if backbone == \"resnet50_radimagenet\":\n        return RADIMAGENET_WEIGHTS_PATH\n    if backbone == \"convnext_base\":\n        return CONVEXT_PATH\n    return EFFICIENTNET_WEIGHTS_PATH\n\n# ================================================================\n# 6. TRAINING (EXACTLY V9 LOGIC)\n# ================================================================\n\ndef ensure_backbone_trained(backbone):\n    cfg = BACKBONE_CONFIG[backbone]\n    epochs, lr = cfg[\"epochs\"], cfg[\"lr\"]\n    scaler = torch.amp.GradScaler(\"cuda\") if USE_AMP else None\n\n    print(f\"\\n{'#' * 70}\\nCHECKING/TRAINING BACKBONE: {backbone}  [{elapsed()}]\\n{'#' * 70}\")\n\n    for fold, (train_idx, val_idx) in enumerate(splits):\n        model_path, meta_path = checkpoint_paths(backbone, fold)\n\n        if os.path.exists(model_path) and os.path.exists(meta_path):\n            print(f\"  Fold {fold+1}/{N_FOLDS}: already trained, skipping. [{elapsed()}]\")\n            continue\n\n        print(f\"\\n  Fold {fold+1}/{N_FOLDS}: MISSING — training now. [{elapsed()}]\")\n\n        fold_train_gold = gold_df.iloc[train_idx].reset_index(drop=True)\n        fold_val = gold_df.iloc[val_idx].reset_index(drop=True)\n        fold_train_df = pd.concat([fold_train_gold, pseudo_df], ignore_index=True)\n\n        train_ds = KneeStudyDataset(fold_train_df, TRAIN_CACHE, training=True)\n        val_ds = KneeStudyDataset(fold_val, TRAIN_CACHE, training=True)\n        train_loader = make_loader(train_ds, shuffle=True)\n        val_loader = make_loader(val_ds, shuffle=False)\n\n        model = AttentionMILModel(\n            backbone_choice=backbone, pretrained_path=pretrained_path_for(backbone)\n        ).to(DEVICE)\n        optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer, max_lr=lr, epochs=epochs, steps_per_epoch=len(train_loader)\n        )\n\n        best_val_auc = -1.0\n        best_state_dict = None\n        best_epoch = -1\n        best_val_preds = None\n\n        for epoch in range(epochs):\n            model.train()\n            running_loss = 0.0\n            pbar = tqdm(train_loader, desc=f\"[{backbone}] F{fold+1} E{epoch+1}/{epochs}\", leave=False)\n            for slices, labels, mask, weight in pbar:\n                slices = slices.to(DEVICE, non_blocking=True)\n                labels = labels.to(DEVICE, non_blocking=True)\n                mask = mask.to(DEVICE, non_blocking=True)\n                weight = weight.to(DEVICE, non_blocking=True)\n\n                optimizer.zero_grad(set_to_none=True)\n                with autocast_context():\n                    logits = model(slices)\n                    loss = masked_weighted_bce(logits, labels, mask, weight, pos_weight)\n\n                if USE_AMP:\n                    scaler.scale(loss).backward()\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    loss.backward()\n                    optimizer.step()\n                scheduler.step()\n                running_loss += loss.item()\n                pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n            model.eval()\n            epoch_val_preds_list = []\n            with torch.no_grad():\n                for slices, labels, mask, weight in val_loader:\n                    slices = slices.to(DEVICE, non_blocking=True)\n                    with autocast_context():\n                        logits = model(slices)\n                    epoch_val_preds_list.append(torch.sigmoid(logits).cpu().numpy())\n            epoch_val_preds = np.concatenate(epoch_val_preds_list, axis=0)\n            epoch_macro_auc, _ = compute_val_macro_auc(fold_val, epoch_val_preds)\n\n            print(f\"  [{backbone}] F{fold+1} E{epoch+1}/{epochs} \"\n                  f\"loss={running_loss/max(len(train_loader),1):.4f} \"\n                  f\"val_auc={epoch_macro_auc:.4f} [{elapsed()}]\")\n\n            if not np.isnan(epoch_macro_auc) and epoch_macro_auc > best_val_auc:\n                best_val_auc = epoch_macro_auc\n                best_epoch = epoch + 1\n                best_state_dict = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n                best_val_preds = epoch_val_preds.copy()\n\n        if best_state_dict is None:\n            best_state_dict = model.state_dict()\n            best_val_preds = epoch_val_preds\n            best_epoch = epochs\n\n        fold_macro_auc, _ = compute_val_macro_auc(fold_val, best_val_preds)\n        torch.save(best_state_dict, model_path)\n        np.savez(meta_path, fold_macro_auc=fold_macro_auc, val_preds=best_val_preds, best_epoch=best_epoch)\n        print(f\"  Fold {fold+1} done: best_epoch={best_epoch} val_auc={fold_macro_auc:.4f} \"\n              f\"saved. [{elapsed()}]\")\n\n        del model, optimizer, scheduler, train_loader, val_loader, train_ds, val_ds, best_state_dict\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    print(f\"\\n{backbone}: all {N_FOLDS} folds present. [{elapsed()}]\")\n\nfor bb in BACKBONE_CONFIG:\n    ensure_backbone_trained(bb)\n\n# ================================================================\n# 7. LOAD OOF PREDICTIONS\n# ================================================================\n\ndef load_oof(backbone):\n    oof = np.full((len(gold_df), N_LABELS), np.nan, dtype=np.float32)\n    for fold, (_, val_idx) in enumerate(splits):\n        _, meta_path = checkpoint_paths(backbone, fold)\n        meta = np.load(meta_path)\n        oof[val_idx] = meta[\"val_preds\"]\n    return oof\n\noof_preds = {name: load_oof(name) for name in BACKBONE_CONFIG}\ngold_true = gold_df[LABEL_COLS].astype(float).values\n\n# ================================================================\n# 8. PER-LABEL ENSEMBLE WEIGHTS\n# ================================================================\n\nprint(f\"\\n{'=' * 70}\\nPER-LABEL COMPARISON  [{elapsed()}]\\n{'=' * 70}\")\nprint(f\"{'Label':20s} \" + \" \".join(f\"{n:>14s}\" for n in BACKBONE_CONFIG) + f\" {'Winner':>12s}\")\n\nlabel_weights = {}\nper_label_aucs = {name: {} for name in BACKBONE_CONFIG}\n\nfor i, col in enumerate(LABEL_COLS):\n    y = gold_true[:, i]\n    valid = ~np.isnan(y)\n    aucs = {}\n    for name in BACKBONE_CONFIG:\n        auc = roc_auc_score(y[valid], oof_preds[name][valid, i]) if len(np.unique(y[valid])) > 1 else 0.5\n        aucs[name] = auc\n        per_label_aucs[name][col] = auc\n\n    winner = max(aucs, key=aucs.get)\n    print(f\"{col:20s} \" + \" \".join(f\"{aucs[n]:14.4f}\" for n in BACKBONE_CONFIG) + f\" {winner:>12s}\")\n\n    temperature = 8.0\n    exp_vals = [np.exp(aucs[n] * temperature) for n in BACKBONE_CONFIG]\n    sum_exp = sum(exp_vals)\n    label_weights[col] = [v / sum_exp for v in exp_vals]\n\nensemble_oof = np.zeros_like(list(oof_preds.values())[0])\nfor i, col in enumerate(LABEL_COLS):\n    w = label_weights[col]\n    for j, name in enumerate(BACKBONE_CONFIG):\n        ensemble_oof[:, i] += w[j] * oof_preds[name][:, i]\n\nprint(f\"\\n{'=' * 70}\")\nfor name in BACKBONE_CONFIG:\n    aucs = [per_label_aucs[name][c] for c in LABEL_COLS if c in per_label_aucs[name]]\n    print(f\"{name:25s} pooled macro-AUC: {np.mean(aucs):.5f}\")\n\nens_aucs = []\nfor i, col in enumerate(LABEL_COLS):\n    y = gold_true[:, i]\n    valid = ~np.isnan(y)\n    if len(np.unique(y[valid])) > 1:\n        ens_aucs.append(roc_auc_score(y[valid], ensemble_oof[valid, i]))\nprint(f\"{'ENSEMBLE':25s} pooled macro-AUC: {np.mean(ens_aucs):.5f}\")\nprint(\"=\" * 70)\n\n# ================================================================\n# 9. TEST INFERENCE (h-flip TTA only)\n# ================================================================\n\ntest_ds = KneeStudyDataset(test, TEST_CACHE, training=False)\ntest_loader = make_loader(test_ds, shuffle=False)\n\ndef run_inference(backbone):\n    print(f\"\\nTest inference: {backbone} [{elapsed()}]\")\n    preds_sum = np.zeros((len(test), N_LABELS), dtype=np.float32)\n    order = []\n    n_models = 0\n\n    for fold in range(N_FOLDS):\n        model_path, _ = checkpoint_paths(backbone, fold)\n        model = AttentionMILModel(\n            backbone_choice=backbone, pretrained_path=pretrained_path_for(backbone)\n        ).to(DEVICE)\n        model.load_state_dict(torch.load(model_path, map_location=DEVICE))\n        model.eval()\n\n        fold_preds = []\n        current_order = []\n        with torch.no_grad():\n            for slices, study_uids in tqdm(test_loader, desc=f\"{backbone} fold {fold+1}\"):\n                slices = slices.to(DEVICE, non_blocking=True)\n                with autocast_context():\n                    logits_a = model(slices)\n                    logits_b = model(torch.flip(slices, dims=[-1]))\n                probs = (torch.sigmoid(logits_a) + torch.sigmoid(logits_b)) / 2.0\n                fold_preds.append(probs.cpu().numpy())\n                current_order.extend(study_uids)\n\n        fold_preds = np.concatenate(fold_preds, axis=0)\n        if n_models == 0:\n            order = current_order\n        preds_sum += fold_preds\n        n_models += 1\n\n        del model\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    return preds_sum / max(n_models, 1), order\n\ntest_preds = {}\nstudy_order = None\nfor name in BACKBONE_CONFIG:\n    test_preds[name], order = run_inference(name)\n    if study_order is None:\n        study_order = order\n\n# ================================================================\n# 10. FINAL BLEND + SUBMISSION\n# ================================================================\n\nfinal_test_preds = np.zeros((len(test), N_LABELS), dtype=np.float32)\nfor i, col in enumerate(LABEL_COLS):\n    w = label_weights[col]\n    for j, name in enumerate(BACKBONE_CONFIG):\n        final_test_preds[:, i] += w[j] * test_preds[name][:, i]\n\nsubmission = pd.DataFrame({\"StudyInstanceUID\": study_order})\nfor i, col in enumerate(LABEL_COLS):\n    submission[col] = final_test_preds[:, i]\nsubmission = submission[sample_sub.columns.tolist()]\n\nsubmission_path = \"/kaggle/working/submission.csv\"\nsubmission.to_csv(submission_path, index=False)\n\nprint(f\"\\n{'=' * 70}\\nSUBMISSION CREATED  [{elapsed()}]\\n{'=' * 70}\")\nprint(f\"Path: {submission_path} | Rows: {len(submission)}\")\nprint(f\"\\nTotal runtime: {elapsed()}\")\nif (time.time() - RUN_START) / 3600 > 8.5:\n    print(\"*** WARNING: runtime close to 9h limit ***\")\nprint()\nprint(submission.head())\nprint(\"\\nDONE\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-26T05:44:24.605456Z","iopub.execute_input":"2026-08-26T05:44:24.606047Z","execution_failed":"2026-08-26T05:45:29.155Z"}},"outputs":[],"execution_count":null}]}