{"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 re\nimport gc\nimport json\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\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\nfrom sklearn.metrics import roc_auc_score\n\ntry:\n    from tqdm.auto import tqdm\nexcept ImportError:\n    tqdm = lambda x, **kwargs: x\n\n\n# ================================================================\n# 1. CONFIG\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\nBASE = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\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\n# ------------------------------------------------\n# SPEED / QUALITY SETTINGS\n# FAST_MODE True  -> quick sanity run (use this FIRST, always)\n# FAST_MODE False -> stronger run for real submissions\n# ------------------------------------------------\nFAST_MODE = True\n\nif FAST_MODE:\n    SLICES_PER_STUDY = 10\n    IMG_SIZE = 224\n    EPOCHS = 6\n    N_FOLDS = 4\n    BACKBONE = \"efficientnet_b0\"\nelse:\n    SLICES_PER_STUDY = 16\n    IMG_SIZE = 256\n    EPOCHS = 15\n    N_FOLDS = 5\n    BACKBONE = \"efficientnet_b3\"\n\nBATCH_SIZE = 4\nLR = 3e-4\nWEIGHT_DECAY = 1e-4\n\nPSEUDO_LABEL_WEIGHT = 0.3     # weight for text-mined labels only\nMIN_GOLD_LABELS_FOR_VAL = 1   # study needs >=this many real labels to be eligible for val fold\nNUM_WORKERS = 2\n\nMAX_SERIES_PER_STUDY = 3      # cap how many series we pull slices from\nSLICES_PER_SERIES_POOL = 6    # candidate slices sampled per chosen series before final pick\n\nCACHE_ROOT = \"/kaggle/working/rsna_knee_cache_v2\"\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\nPRETRAINED_WEIGHTS_PATH = (\n    \"/kaggle/input/models/elizavetanew/\"\n    \"efficientnet_b0/pytorch/\"\n    \"efficientnet_b0_rwightman-7f5810bc/1/\"\n    \"efficientnet_b0_rwightman-7f5810bc.pth\"\n)\n\nprint(\"=\" * 70)\nprint(\"RSNA KNEE ABNORMALITY DETECTION — V2\")\nprint(\"=\" * 70)\nprint(f\"Device: {DEVICE} | FAST_MODE: {FAST_MODE} | Backbone: {BACKBONE}\")\nprint(f\"Slices/study: {SLICES_PER_STUDY} | Img size: {IMG_SIZE}\")\nprint(f\"Epochs: {EPOCHS} | Folds: {N_FOLDS} | Batch: {BATCH_SIZE}\")\nprint(\"=\" * 70)\n\n\n# ================================================================\n# 2. LOAD DATA (defensive — inspect what actually exists)\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\nprint(\"\\ntrain.csv columns:\", train.columns.tolist())\nprint(\"train_series.csv columns:\", train_series.columns.tolist())\n\nREPORT_COL = None\nfor cand in [\"Report\", \"report\", \"ReportText\", \"report_text\", \"Findings\"]:\n    if cand in train.columns:\n        REPORT_COL = cand\n        break\nif REPORT_COL is None:\n    print(\"WARNING: no report/text column found — weak-label mining disabled.\")\n\n# True per-label gold coverage (this is the number that actually matters)\nprint(\"\\nReal (non-null) label coverage per column:\")\ngold_counts = {}\nfor col in LABEL_COLS:\n    if col in train.columns:\n        n = train[col].notna().sum()\n    else:\n        n = 0\n        train[col] = np.nan\n    gold_counts[col] = n\n    print(f\"  {col:20s}: {n} / {len(train)}\")\n\n# a study counts as \"has any gold\" if ANY label column is non-null for it\nany_gold_mask = train[LABEL_COLS].notna().any(axis=1)\nprint(f\"\\nStudies with >=1 real label in ANY column: {any_gold_mask.sum()} / {len(train)}\")\nprint(f\"Test studies: {len(test)}\")\n\n# Try to find a patient/subject id for GroupKFold (avoid leakage if\n# the same patient appears in multiple studies)\nGROUP_COL = None\nfor cand in [\"PatientID\", \"patient_id\", \"SubjectID\", \"subject_id\"]:\n    if cand in train.columns:\n        GROUP_COL = cand\n        break\nprint(f\"Group column for CV: {GROUP_COL if GROUP_COL else 'none found (will use plain K-fold / stratified)'}\")\n\n\n# ================================================================\n# 3. WEAK-LABEL MINING (only if report text exists)\n# ================================================================\n\nNEGATION_WORDS = [\n    \"no\", \"without\", \"not\", \"negative for\", \"no evidence of\",\n    \"no signs of\", \"intact\", \"unremarkable\", \"no significant\",\n    \"absence of\", \"no hay\", \"sin\", \"ausencia de\", \"no se observa\",\n    \"intacto\", \"intacta\", \"no evidencia\", \"normal\", \"kein\", \"keine\",\n    \"ohne\", \"unauffällig\", \"intakt\", \"kein hinweis\",\n]\n\nLABEL_KEYWORDS = {\n    \"ACL\": [\"acl\", \"anterior cruciate\", \"ligamento cruzado anterior\", \"lca\",\n            \"vorderes kreuzband\", \"kreuzbandruptur\"],\n    \"MCL\": [\"mcl\", \"medial collateral\", \"ligamento colateral medial\",\n            \"mediales seitenband\"],\n    \"Medial Meniscus\": [\"medial meniscus\", \"menisco medial\", \"medialer meniskus\"],\n    \"Lateral Meniscus\": [\"lateral meniscus\", \"menisco lateral\", \"lateraler meniskus\"],\n    \"Medial OA\": [\"medial compartment osteoarthritis\", \"medial osteoarthritis\",\n                  \"artrosis medial\", \"medialarthrose\", \"medial compartment degenerat\",\n                  \"medial compartment\", \"medial joint space narrowing\",\n                  \"medial chondromalacia\", \"medial cartilage loss\"],\n    \"Lateral OA\": [\"lateral compartment osteoarthritis\", \"lateral osteoarthritis\",\n                   \"artrosis lateral\", \"lateralarthrose\", \"lateral compartment\",\n                   \"lateral joint space narrowing\", \"lateral chondromalacia\",\n                   \"lateral cartilage loss\"],\n    \"PF OA\": [\"patellofemoral osteoarthritis\", \"patellofemoral joint osteoarthritis\",\n              \"artrosis patelofemoral\", \"femoropatellararthrose\", \"patellofemoral\",\n              \"patellofemoral joint space narrowing\", \"patellofemoral chondromalacia\",\n              \"patellar cartilage loss\"],\n    \"Effusion\": [\"effusion\", \"joint effusion\", \"derrame articular\", \"erguss\",\n                 \"gelenkerguss\"],\n    \"Synovitis\": [\"synovitis\", \"sinovitis\", \"synovialitis\"],\n    \"Baker's\": [\"baker's cyst\", \"baker cyst\", \"quiste de baker\", \"popliteal cyst\",\n                \"bakerzyste\"],\n    \"Contusion\": [\"contusion\", \"bone bruise\", \"contusión ósea\", \"contusion osea\",\n                  \"knochenkontusion\", \"bone marrow edema\"],\n    \"Fracture\": [\"fracture\", \"fx\", \"fractura\", \"fraktur\"],\n}\n\n\ndef mine_label(report_text, keywords):\n    if not isinstance(report_text, str) or not report_text.strip():\n        return None\n    text = report_text.lower()\n    clauses = re.split(r\"[.\\n;]\", text)\n    found_any = False\n    for clause in clauses:\n        for kw in keywords:\n            if kw in clause:\n                found_any = True\n                idx = clause.find(kw)\n                before = clause[max(0, idx - 40):idx]\n                after = clause[idx + len(kw): idx + len(kw) + 25]\n                is_negated = (\n                    any(neg in before for neg in NEGATION_WORDS)\n                    or any(neg in after for neg in NEGATION_WORDS)\n                )\n                if not is_negated:\n                    return 1\n    return 0 if found_any else None\n\n\ndef mine_all_labels(report_text):\n    return {col: mine_label(report_text, LABEL_KEYWORDS[col]) for col in LABEL_COLS}\n\n\nif REPORT_COL is not None:\n    print(\"\\nMining report labels...\")\n    mined = train[REPORT_COL].apply(mine_all_labels).apply(pd.Series)\n    mined.columns = [f\"{c}_mined\" for c in LABEL_COLS]\n    train = pd.concat([train, mined], axis=1)\nelse:\n    for col in LABEL_COLS:\n        train[f\"{col}_mined\"] = np.nan\n\n\n# ================================================================\n# 4. VALIDATE MINED LABELS AGAINST REAL LABELS (per-column, using\n#    ALL rows that have a real value for that column — not just the\n#    58-study first-column mask from V1)\n# ================================================================\n\nprint(\"\\nWeak-label validation against real labels (per-column, all available):\")\nTRUSTED_MINED_COLS = []\nAGREEMENT_THRESHOLD = 0.65\nMIN_OVERLAP_N = 15  # don't trust an agreement stat computed on tiny n\n\nfor col in LABEL_COLS:\n    real = train[col]\n    pred = train[f\"{col}_mined\"]\n    valid = real.notna() & pred.notna()\n    n = valid.sum()\n    if n >= MIN_OVERLAP_N:\n        acc = (real[valid] == pred[valid]).mean()\n        trusted = acc >= AGREEMENT_THRESHOLD\n        if trusted:\n            TRUSTED_MINED_COLS.append(col)\n        status = \"USED\" if trusted else \"DROPPED\"\n        print(f\"  {col:20s} agreement={acc:.3f} n={n} [{status}]\")\n    else:\n        print(f\"  {col:20s} insufficient overlap n={n} [DROPPED]\")\n\nfor col in LABEL_COLS:\n    if col not in TRUSTED_MINED_COLS:\n        train[f\"{col}_mined\"] = np.nan\n\nprint(\"\\nTrusted pseudo-label columns:\", TRUSTED_MINED_COLS)\n\n\n# ================================================================\n# 5. FINAL LABEL + ELEMENT-WISE WEIGHT MATRIX\n#    (this is the core fix — weight is per (study,label), not per study)\n# ================================================================\n\nfinal_labels = train[[\"StudyInstanceUID\"]].copy()\nweight_matrix = np.zeros((len(train), N_LABELS), dtype=np.float32)\nis_real_matrix = np.zeros((len(train), N_LABELS), dtype=bool)\n\nfor i, col in enumerate(LABEL_COLS):\n    real_col = train[col]\n    mined_col = train[f\"{col}_mined\"]\n\n    combined = real_col.where(real_col.notna(), mined_col)\n    final_labels[col] = combined\n\n    real_mask = real_col.notna().values\n    mined_only_mask = (~real_mask) & mined_col.notna().values\n\n    weight_matrix[real_mask, i] = 1.0\n    weight_matrix[mined_only_mask, i] = PSEUDO_LABEL_WEIGHT\n    is_real_matrix[real_mask, i] = True\n\nfinal_labels[\"n_real_labels\"] = is_real_matrix.sum(axis=1)\nfinal_labels[\"n_any_labels\"] = final_labels[LABEL_COLS].notna().sum(axis=1).values\n\nfor i, col in enumerate(LABEL_COLS):\n    final_labels[f\"__w_{col}\"] = weight_matrix[:, i]\n\nhas_any_label = final_labels[\"n_any_labels\"] > 0\nfinal_labels = final_labels[has_any_label].reset_index(drop=True)\n\nif GROUP_COL is not None:\n    final_labels[GROUP_COL] = train.loc[has_any_label.values, GROUP_COL].values\n\nprint(f\"\\nUsable training studies: {len(final_labels)}\")\nprint(f\"Studies with >=1 REAL label: {(final_labels['n_real_labels'] > 0).sum()}\")\n\n\n# ================================================================\n# 6. SERIES LOOKUP + PLANE/SEQUENCE SCORING\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        lookup.setdefault(uid, []).append(str(row.SeriesInstanceUID))\n    return lookup\n\n\nprint(\"\\nBuilding series lookup...\")\nTRAIN_SERIES_LOOKUP = build_series_lookup(train_series)\nTEST_SERIES_LOOKUP = build_series_lookup(test_series)\nprint(f\"Train studies in lookup: {len(TRAIN_SERIES_LOOKUP)}\")\nprint(f\"Test studies in lookup: {len(TEST_SERIES_LOOKUP)}\")\n\nSCOUT_KEYWORDS = [\"scout\", \"localizer\", \"loc\", \"survey\", \"calibration\"]\nGOOD_PLANE_KEYWORDS = [\"sag\", \"cor\", \"ax\", \"t1\", \"t2\", \"pd\", \"stir\",\n                        \"fat sat\", \"fs\", \"prot\", \"dess\"]\n\n\ndef score_series_by_description(series_dir):\n    \"\"\"Peek at first DICOM in a series to score how 'diagnostic' it looks.\n    Higher score = more likely a real diagnostic sequence, not a scout.\"\"\"\n    try:\n        files = [f for f in os.listdir(series_dir) if not f.startswith(\".\")]\n        if not files:\n            return -1.0\n        files.sort()\n        sample_path = os.path.join(series_dir, files[len(files) // 2])\n        dcm = pydicom.dcmread(sample_path, stop_before_pixels=False)\n    except Exception:\n        return -1.0\n\n    desc = str(getattr(dcm, \"SeriesDescription\", \"\")).lower()\n    score = 0.0\n\n    if any(k in desc for k in SCOUT_KEYWORDS):\n        score -= 5.0\n    if any(k in desc for k in GOOD_PLANE_KEYWORDS):\n        score += 2.0\n\n    # Prefer series with a reasonable number of slices (real volumes,\n    # not single-shot scouts which often have very few images)\n    try:\n        n_files = len(files)\n        if n_files >= 10:\n            score += 1.0\n        elif n_files <= 2:\n            score -= 2.0\n    except Exception:\n        pass\n\n    # Prefer typical in-plane resolution for diagnostic MRI (avoid tiny\n    # thumbnails / oddly shaped scouts)\n    try:\n        rows = int(getattr(dcm, \"Rows\", 0))\n        cols = int(getattr(dcm, \"Columns\", 0))\n        if rows >= 200 and cols >= 200:\n            score += 0.5\n    except Exception:\n        pass\n\n    return score\n\n\ndef pick_best_series(study_dir, series_ids, max_series):\n    scored = []\n    for series_uid in series_ids:\n        series_dir = os.path.join(study_dir, series_uid)\n        if not os.path.isdir(series_dir):\n            continue\n        s = score_series_by_description(series_dir)\n        scored.append((s, series_uid, series_dir))\n\n    if not scored:\n        return []\n\n    scored.sort(key=lambda x: x[0], reverse=True)\n    return [(uid, path) for _, uid, path in scored[:max_series]]\n\n\n# ================================================================\n# 7. FIND DICOM SLICES (plane-aware)\n# ================================================================\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_ids = series_lookup.get(study_uid, [])\n\n    chosen_series = pick_best_series(study_dir, series_ids, MAX_SERIES_PER_STUDY)\n\n    all_slices = []\n    for series_uid, series_dir in chosen_series:\n        try:\n            files = [f for f in os.listdir(series_dir) if not f.startswith(\".\")]\n        except Exception:\n            continue\n        files.sort()\n        if not files:\n            continue\n\n        count = min(SLICES_PER_SERIES_POOL, len(files))\n        # Bias toward the middle of the stack — central slices of a knee\n        # MRI volume are far more likely to show the joint itself than\n        # the first/last few slices near the skin surface.\n        center = len(files) / 2.0\n        spread = len(files) * 0.35\n        lo = max(0, int(center - spread))\n        hi = min(len(files) - 1, int(center + spread))\n        if hi <= lo:\n            lo, hi = 0, len(files) - 1\n        indices = np.linspace(lo, hi, count).astype(int)\n\n        for idx in indices:\n            all_slices.append(os.path.join(series_dir, files[idx]))\n\n    if not all_slices:\n        return []\n\n    if len(all_slices) >= n_slices:\n        pick_indices = np.linspace(0, len(all_slices) - 1, n_slices).astype(int)\n        chosen = [all_slices[i] for i in pick_indices]\n    else:\n        chosen = all_slices + [all_slices[-1]] * (n_slices - len(all_slices))\n\n    return chosen\n\n\n# ================================================================\n# 8. DICOM -> IMAGE\n# ================================================================\n\ndef process_dicom(path):\n    try:\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n\n        # apply rescale slope/intercept if present (real intensity units)\n        slope = float(getattr(dcm, \"RescaleSlope\", 1.0) or 1.0)\n        intercept = float(getattr(dcm, \"RescaleIntercept\", 0.0) or 0.0)\n        img = img * slope + intercept\n\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\n\n# ================================================================\n# 9. CACHE\n# ================================================================\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\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\n\ndef build_cache(study_uids, series_lookup, split, cache_dir, name):\n    print(f\"\\nBuilding {name} cache...\")\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}\")\n\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# ================================================================\n# 10. DATASET\n# ================================================================\n\nW_COLS = [f\"__w_{c}\" for c in LABEL_COLS]\n\n\nclass KneeStudyDataset(Dataset):\n    def __init__(self, df, cache_dir, training=True, augment=False):\n        self.df = df.reset_index(drop=True).copy()\n        self.cache_dir = cache_dir\n        self.training = training\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.df)\n\n    def _augment(self, slices):\n        # slices: (S, H, W)\n        if random.random() < 0.5:\n            slices = slices[:, :, ::-1].copy()\n        if random.random() < 0.3:\n            k = random.choice([1, 2, 3])\n            slices = np.rot90(slices, k=k, axes=(1, 2)).copy()\n        return slices\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        uid = str(row[\"StudyInstanceUID\"])\n        cache_path = os.path.join(self.cache_dir, f\"{uid}.npy\")\n        slices = np.load(cache_path).astype(np.float32)\n\n        if not self.training:\n            slices_t = torch.from_numpy(slices.copy()).unsqueeze(1)\n            return slices_t, uid\n\n        if self.augment:\n            slices = self._augment(slices)\n\n        label_values = row[LABEL_COLS].to_numpy(dtype=object)\n        weight_values = row[W_COLS].to_numpy(dtype=np.float32)\n\n        labels = np.zeros(N_LABELS, dtype=np.float32)\n        label_mask = np.zeros(N_LABELS, dtype=np.float32)\n\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\n        return (\n            torch.from_numpy(slices.copy()).unsqueeze(1),\n            torch.from_numpy(labels),\n            torch.from_numpy(label_mask),\n            torch.from_numpy(weight_values.copy()),\n        )\n\n\n# ================================================================\n# 11. MODEL\n# ================================================================\n\nclass AttentionMILModel(nn.Module):\n    def __init__(self, n_labels=N_LABELS, pretrained_path=None, backbone_name=\"efficientnet_b0\"):\n        super().__init__()\n        import torchvision.models as tvm\n\n        if backbone_name == \"efficientnet_b3\":\n            backbone = tvm.efficientnet_b3(weights=None)\n            feat_dim = 1536\n        else:\n            backbone = tvm.efficientnet_b0(weights=None)\n            feat_dim = 1280\n\n        if pretrained_path is not None and os.path.exists(pretrained_path) and backbone_name == \"efficientnet_b0\":\n            print(\"Loading offline pretrained weights...\")\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            missing, unexpected = backbone.load_state_dict(cleaned, strict=False)\n            print(f\"Pretrained weights loaded. missing={len(missing)}, unexpected={len(unexpected)}\")\n        else:\n            print(\"WARNING: pretrained weights not applied (path missing or backbone mismatch). Training from scratch/ImageNet default only if available.\")\n\n        self.encoder = backbone.features\n        self.pool = nn.AdaptiveAvgPool2d(1)\n\n        self.attention = nn.Sequential(\n            nn.Linear(feat_dim, 128), nn.Tanh(), nn.Linear(128, 1)\n        )\n\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)\n        x = x.repeat(1, 3, 1, 1)\n        feats = self.encoder(x)\n        feats = self.pool(feats).flatten(1)\n        feats = feats.reshape(B, S, -1)\n        attn_scores = self.attention(feats)\n        attn_weights = torch.softmax(attn_scores, dim=1)\n        study_feat = (feats * attn_weights).sum(dim=1)\n        return self.head(study_feat)\n\n\n# ================================================================\n# 12. LOSS (per-element weight matrix, not per-study scalar)\n# ================================================================\n\ndef masked_weighted_bce(logits, labels, label_mask, weight_matrix, pos_weight):\n    loss = F.binary_cross_entropy_with_logits(\n        logits, labels, pos_weight=pos_weight, reduction=\"none\"\n    )\n    loss = loss * label_mask * weight_matrix\n    denom = (label_mask * weight_matrix).sum(dim=1).clamp(min=1e-6)\n    loss = loss.sum(dim=1) / denom\n    return loss.mean()\n\n\n# ================================================================\n# 13. POSITIVE CLASS WEIGHTS (computed on REAL labels only, capped)\n# ================================================================\n\nreal_label_block = train[LABEL_COLS].astype(float)\npos_counts = real_label_block.sum()\nneg_counts = real_label_block.apply(lambda c: (c == 0).sum())\npos_weight_values = (neg_counts / pos_counts.clip(lower=1)).values\npos_weight_values = np.clip(pos_weight_values, 0.2, 8.0)  # avoid exploding on rare/tiny-n labels\npos_weight = torch.tensor(pos_weight_values, dtype=torch.float32).to(DEVICE)\n\nprint(\"\\nPositive weights (capped 0.2-8.0):\")\nfor col, value in zip(LABEL_COLS, pos_weight_values):\n    print(f\"  {col:20s}: {value:.3f}\")\n\n\n# ================================================================\n# 14. CV SPLITS\n#    Prefer GroupKFold (by patient) if we have a group col, to avoid\n#    leakage. Fall back to StratifiedKFold on a composite label\n#    (mean positivity across labels, binned) else plain KFold.\n# ================================================================\n\nn = len(final_labels)\nN_FOLDS_EFFECTIVE = min(N_FOLDS, max(2, n // 20))\nif N_FOLDS_EFFECTIVE != N_FOLDS:\n    print(f\"\\nReducing folds from {N_FOLDS} to {N_FOLDS_EFFECTIVE} given dataset size.\")\nN_FOLDS = N_FOLDS_EFFECTIVE\n\nif GROUP_COL is not None and GROUP_COL in final_labels.columns:\n    gkf = GroupKFold(n_splits=N_FOLDS)\n    splits = list(gkf.split(final_labels, groups=final_labels[GROUP_COL]))\n    print(f\"Using GroupKFold on '{GROUP_COL}' with {N_FOLDS} folds.\")\nelse:\n    composite = final_labels[LABEL_COLS].astype(float).mean(axis=1, skipna=True).fillna(0)\n    strat_bins = pd.qcut(composite.rank(method=\"first\"), q=min(N_FOLDS, 10), labels=False, duplicates=\"drop\")\n    try:\n        skf = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n        splits = list(skf.split(final_labels, strat_bins))\n        print(f\"Using StratifiedKFold on composite label bins with {N_FOLDS} folds.\")\n    except Exception:\n        kf = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n        splits = list(kf.split(final_labels))\n        print(f\"Falling back to plain KFold with {N_FOLDS} folds.\")\n\nfold_aucs = []\noof_preds = np.zeros((n, N_LABELS), dtype=np.float32)\n\nUSE_AMP = DEVICE.type == \"cuda\"\nscaler = torch.amp.GradScaler(\"cuda\") if USE_AMP else None\n\n\ndef autocast_context():\n    if USE_AMP:\n        return torch.amp.autocast(\"cuda\")\n    return torch.autocast(device_type=\"cpu\", enabled=False)\n\n\ndef make_loader(dataset, shuffle):\n    kwargs = dict(\n        batch_size=BATCH_SIZE, shuffle=shuffle, num_workers=NUM_WORKERS,\n        pin_memory=(DEVICE.type == \"cuda\"),\n    )\n    if NUM_WORKERS > 0:\n        kwargs[\"persistent_workers\"] = True\n    return DataLoader(dataset, **kwargs)\n\n\n# ================================================================\n# 15. TRAINING\n# ================================================================\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING CROSS VALIDATION\")\nprint(\"=\" * 70)\n\nfor fold, (train_idx, val_idx) in enumerate(splits):\n    print(f\"\\n{'=' * 70}\\nFOLD {fold + 1}/{N_FOLDS}\\n{'=' * 70}\")\n\n    fold_train_df = final_labels.iloc[train_idx].reset_index(drop=True)\n    fold_val_df = final_labels.iloc[val_idx].reset_index(drop=True)\n\n    print(f\"Train studies: {len(fold_train_df)} | Validation studies: {len(fold_val_df)}\")\n\n    train_ds = KneeStudyDataset(fold_train_df, TRAIN_CACHE, training=True, augment=True)\n    val_ds = KneeStudyDataset(fold_val_df, TRAIN_CACHE, training=True, augment=False)\n\n    train_loader = make_loader(train_ds, shuffle=True)\n    val_loader = make_loader(val_ds, shuffle=False)\n\n    model = AttentionMILModel(\n        pretrained_path=PRETRAINED_WEIGHTS_PATH, backbone_name=BACKBONE\n    ).to(DEVICE)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=LR, epochs=EPOCHS, steps_per_epoch=max(len(train_loader), 1)\n    )\n\n    for epoch in range(EPOCHS):\n        model.train()\n        running_loss = 0.0\n        pbar = tqdm(train_loader, desc=f\"Fold {fold+1} Epoch {epoch+1}/{EPOCHS}\", leave=True)\n\n        for slices, labels, mask, wmat 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            wmat = wmat.to(DEVICE, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n\n            with autocast_context():\n                logits = model(slices)\n                loss = masked_weighted_bce(logits, labels, mask, wmat, pos_weight)\n\n            if USE_AMP:\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n                optimizer.step()\n\n            scheduler.step()\n            running_loss += loss.item()\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n        print(f\"Epoch {epoch+1}/{EPOCHS} train_loss={running_loss / max(len(train_loader),1):.5f}\")\n\n    # ---------------- VALIDATION ----------------\n    model.eval()\n    val_preds, val_targets = [], []\n\n    with torch.no_grad():\n        for slices, labels, mask, wmat in val_loader:\n            slices = slices.to(DEVICE, non_blocking=True)\n            with autocast_context():\n                logits = model(slices)\n            probs = torch.sigmoid(logits)\n            val_preds.append(probs.cpu().numpy())\n            val_targets.append(labels.numpy())\n\n    val_preds = np.concatenate(val_preds, axis=0)\n    val_targets = np.concatenate(val_targets, axis=0)\n    oof_preds[val_idx] = val_preds\n\n    aucs = []\n    print(\"\\nValidation AUC (computed only where a REAL label exists):\")\n    for i, col in enumerate(LABEL_COLS):\n        real_vals = train.set_index(\"StudyInstanceUID\").reindex(\n            fold_val_df[\"StudyInstanceUID\"]\n        )[col].astype(float).values\n        valid = ~np.isnan(real_vals)\n        if valid.sum() > 5 and len(np.unique(real_vals[valid])) > 1:\n            auc = roc_auc_score(real_vals[valid], val_preds[valid, i])\n            aucs.append(auc)\n            print(f\"  {col:20s}: {auc:.4f} (n={valid.sum()})\")\n        else:\n            print(f\"  {col:20s}: skipped (n={valid.sum()}, insufficient real labels)\")\n\n    fold_auc = np.mean(aucs) if aucs else float(\"nan\")\n    fold_aucs.append(fold_auc)\n    print(f\"\\nFold {fold+1} Macro-AUC (on real labels only): {fold_auc:.5f}\")\n\n    model_path = f\"/kaggle/working/model_fold{fold}.pt\"\n    torch.save(model.state_dict(), model_path)\n    print(f\"Saved: {model_path}\")\n\n    del model, optimizer, scheduler, train_loader, val_loader, train_ds, val_ds\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\n# ================================================================\n# 16. CV RESULT\n# ================================================================\n\nmean_cv = np.nanmean(fold_aucs)\nstd_cv = np.nanstd(fold_aucs)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CROSS-VALIDATION COMPLETE\")\nprint(\"=\" * 70)\nfor i, score in enumerate(fold_aucs):\n    print(f\"Fold {i+1}: {score:.5f}\")\nprint(f\"\\nCV Macro-AUC: {mean_cv:.5f} +/- {std_cv:.5f}\")\n\nif mean_cv < 0.65:\n    print(\"\\n*** WARNING: CV is still weak. Before submitting, check: ***\")\n    print(\"  - How many REAL (non-mined) labels actually exist per column (see section 2 printout)\")\n    print(\"  - Whether MAX_SERIES_PER_STUDY / plane scoring picked reasonable series\")\n    print(\"  - Consider FAST_MODE=False for a stronger final run once CV logic looks right\")\n\n\n# ================================================================\n# 17. TEST INFERENCE (fold ensemble + horizontal-flip TTA)\n# ================================================================\n\ntest_ds = KneeStudyDataset(test, TEST_CACHE, training=False)\ntest_loader = make_loader(test_ds, shuffle=False)\n\nall_preds = np.zeros((len(test), N_LABELS), dtype=np.float32)\nstudy_order = []\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TEST INFERENCE\")\nprint(\"=\" * 70)\n\nfor fold in range(N_FOLDS):\n    print(f\"\\nLoading fold {fold+1}/{N_FOLDS}\")\n    model = AttentionMILModel(\n        pretrained_path=PRETRAINED_WEIGHTS_PATH, backbone_name=BACKBONE\n    ).to(DEVICE)\n    model.load_state_dict(torch.load(f\"/kaggle/working/model_fold{fold}.pt\", map_location=DEVICE))\n    model.eval()\n\n    fold_preds = []\n    current_order = []\n\n    with torch.no_grad():\n        for slices, study_uids in tqdm(test_loader, desc=f\"Inference fold {fold+1}\"):\n            slices = slices.to(DEVICE, non_blocking=True)\n\n            with autocast_context():\n                logits_a = model(slices)\n                logits_b = model(torch.flip(slices, dims=[-1]))  # horizontal flip TTA\n\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\n    if fold == 0:\n        study_order = current_order\n    else:\n        if current_order != study_order:\n            raise RuntimeError(\"Test study ordering changed between folds.\")\n\n    all_preds += fold_preds / N_FOLDS\n\n    del model\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n\n# ================================================================\n# 18. CREATE SUBMISSION\n# ================================================================\n\nsubmission = pd.DataFrame({\"StudyInstanceUID\": study_order})\nfor i, col in enumerate(LABEL_COLS):\n    submission[col] = all_preds[:, i]\n\nsubmission = submission[sample_sub.columns.tolist()]\nsubmission_path = \"/kaggle/working/submission.csv\"\nsubmission.to_csv(submission_path, index=False)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SUBMISSION CREATED\")\nprint(\"=\" * 70)\nprint(f\"Path: {submission_path}\")\nprint(f\"Rows: {len(submission)} | Columns: {len(submission.columns)}\")\nprint(\"\\nPrediction ranges:\")\nfor col in LABEL_COLS:\n    print(f\"  {col:20s}: {submission[col].min():.4f} -> {submission[col].max():.4f}\")\nprint()\nprint(submission.head())\nprint(\"\\n\" + \"=\" * 70)\nprint(\"DONE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-13T13:18:14.017293Z","iopub.execute_input":"2026-08-13T13:18:14.017509Z","iopub.status.idle":"2026-08-13T13:48:12.612115Z","shell.execute_reply.started":"2026-08-13T13:18:14.017488Z","shell.execute_reply":"2026-08-13T13:48:12.611035Z"}},"outputs":[],"execution_count":null}]}