{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA Knee Abnormality Detection — Dual-Sequence Multi-Level DINOv2 v6.1\n\nA stronger but still practical upgrade over the 0.621 multi-plane baseline.\n\nMain changes:\n\n- six modality slots: fluid-sensitive and structural series for each plane;\n- one wide-context three-slice RGB representation per selected series;\n- multi-level DINOv2 statistics from intermediate and final blocks;\n- sequence metadata and global study-protocol metadata;\n- target-conditioned attention over six sequence slots;\n- stochastic modality dropout and feature noise;\n- target-to-target self-attention for clinically correlated findings;\n- report teacher enlarged with high-confidence rule-derived examples;\n- honest exact-label OOF comparison of weak-supervision and exact-only heads;\n- per-target OOF-derived blending and rank ensemble.\n\nThe image pipeline remains frozen-feature based and uses both T4 GPUs.","metadata":{}},{"cell_type":"code","source":"import os, gc, re, json, math, time, random, warnings, unicodedata, zlib\nfrom copy import deepcopy\nfrom pathlib import Path\nfrom collections import defaultdict\nfrom concurrent.futures import ThreadPoolExecutor\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nfrom sklearn.pipeline import FeatureUnion\nfrom sklearn.feature_extraction.text import TfidfVectorizer\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score, balanced_accuracy_score\nfrom scipy.stats import rankdata\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\nimport pydicom\n\nwarnings.filterwarnings(\"ignore\")\ntorch.set_float32_matmul_precision(\"high\")\nif torch.cuda.is_available():\n    torch.backends.cudnn.benchmark = True\n\n\nclass CFG:\n    seed = 2026\n    debug = False\n\n    # Strong image profile\n    img_size = 126\n    n_slices = 1\n    n_slots = 6\n    center_header_samples = 11\n    dino_layers = 12  # Use all 12 blocks for maximal feature richness\n    feature_batch_size = 256\n    feature_loader_batch_size = 48\n    feature_workers = 8\n    rebuild_features = False\n    feature_cache_tag = \"dual_sequence_wide2p5d_dino126_l12_multilevel_v2\"\n\n    # Multimodal target head\n    hidden_dim = 320\n    dropout = 0.18\n    modality_dropout = 0.16\n    feature_noise = 0.012\n    train_batch_size = 256\n    pretrain_epochs = 8  # Slight increase to stabilize pretraining\n    steps_per_epoch = 24\n    finetune_epochs = 6\n    finetune_steps = 12\n    exact_only_epochs = 18\n    exact_only_steps = 9\n    lr = 6e-4\n    finetune_lr = 1.2e-4\n    exact_only_lr = 2.2e-4\n    weight_decay = 3e-4\n    rank_loss_weight = 0.07\n    finetune_rank_weight = 0.20\n    exact_only_rank_weight = 0.28\n    exact_sample_boost = 22.0\n    ema_decay = 0.995\n\n    # OOF branch selection\n    run_oof_blending = True\n    oof_folds = 4\n    final_weak_seeds = (2026, 2037, 2051, 2067)  # Added 4th seed to lower ensemble variance\n    final_exact_seeds = (2081, 2099, 2111)       # Added 3rd seed to exact ensemble\n\n    # Reliability-aware report supervision\n    rule_threshold = 0.10\n    text_threshold = 0.30\n    min_pseudo_weight = 0.12\n    max_pseudo_weight = 0.82\n    pseudo_global_scale = 0.72\n    rule_teacher_weight = 0.22\n    compute_text_oof = True\n\n    competition_dir = None\n    dinov2_dir = None\n\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    amp = torch.cuda.is_available()\n    n_gpus = torch.cuda.device_count()\n\n\nTARGETS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\", \"Effusion\",\n    \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\"\n]\n\nPLANES = [\"Sagittal\", \"Coronal\", \"Axial\"]\nCONTRASTS = [\"Fluid\", \"Structural\"]\nSLOT_NAMES = [\n    f\"{plane}_{contrast}\"\n    for plane in PLANES\n    for contrast in CONTRASTS\n]\nSLOT_META_DIM = 4\nSTUDY_META_DIM = 12\n\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n\n\ndef find_competition_dir():\n    if CFG.competition_dir is not None:\n        p = Path(CFG.competition_dir)\n        if p.exists():\n            return p\n\n    expected = Path(\"/kaggle/input/competitions/rsna-knee-abnormality-detection\")\n    if expected.exists():\n        return expected\n\n    for root in (Path(\"/kaggle/input\"), Path(\".\")):\n        if not root.exists():\n            continue\n        for train_csv in root.glob(\"*/train.csv\"):\n            p = train_csv.parent\n            required = [\n                p / \"train_series.csv\", p / \"test.csv\", p / \"test_series.csv\",\n                p / \"train_series\", p / \"test_series\", p / \"sample_submission.csv\"\n            ]\n            if all(x.exists() for x in required):\n                return p\n\n    raise FileNotFoundError(\"Competition directory was not found.\")\n\n\ndef find_dinov2_dir():\n    if CFG.dinov2_dir is not None:\n        p = Path(CFG.dinov2_dir)\n        if (p / \"config.json\").exists():\n            return p\n\n    search_root = Path(\"/kaggle/input\")\n    skip_dirs = {\"train_series\", \"test_series\", \".git\", \"__pycache__\"}\n    candidates = []\n\n    if search_root.exists():\n        base_depth = len(search_root.parts)\n        for root, dirs, files in os.walk(search_root):\n            depth = len(Path(root).parts) - base_depth\n            dirs[:] = [d for d in dirs if d not in skip_dirs and depth < 7]\n            if \"config.json\" not in files:\n                continue\n\n            root_path = Path(root)\n            has_weights = any(\n                (root_path / name).exists()\n                for name in (\"model.safetensors\", \"pytorch_model.bin\")\n            )\n            if not has_weights:\n                continue\n\n            try:\n                cfg = json.loads((root_path / \"config.json\").read_text())\n            except Exception:\n                continue\n\n            model_type = str(cfg.get(\"model_type\", \"\")).lower()\n            name_hint = str(root_path).lower()\n            if \"dinov2\" in model_type or \"dinov2\" in name_hint:\n                hidden = int(cfg.get(\"hidden_size\", 10_000))\n                candidates.append((abs(hidden - 384), root_path))\n\n    if not candidates:\n        raise FileNotFoundError(\n            \"DINOv2 weights were not found. Add facebook/dinov2-small as a Kaggle model input \"\n            \"or set CFG.dinov2_dir to its local folder.\"\n        )\n\n    return sorted(candidates, key=lambda x: x[0])[0][1]\n\n\nseed_everything(CFG.seed)\nBASE = find_competition_dir()\nWORK = Path(\"/kaggle/working\")\nWORK.mkdir(parents=True, exist_ok=True)\n\nprint(\"Device:\", CFG.device)\nprint(\"Visible GPUs:\", CFG.n_gpus, [torch.cuda.get_device_name(i) for i in range(CFG.n_gpus)])\nprint(\"Competition:\", BASE)\nprint(\n    \"Profile:\",\n    f\"{CFG.n_slots} plane/sequence slots × 3 sampled positions,\",\n    f\"DINO blocks={CFG.dino_layers}, image={CFG.img_size}\"\n)","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(BASE / \"train.csv\")\ntrain_series = pd.read_csv(BASE / \"train_series.csv\")\ntest = pd.read_csv(BASE / \"test.csv\")\ntest_series = pd.read_csv(BASE / \"test_series.csv\")\nsample_submission = pd.read_csv(BASE / \"sample_submission.csv\")\n\n\ndef normalize_column_names(df):\n    df = df.copy()\n    df.columns = [str(col).replace(\"\\ufeff\", \"\").strip() for col in df.columns]\n\n    canonical = {\n        \"studyinstanceuid\": \"StudyInstanceUID\",\n        \"patientsex\": \"PatientSex\",\n        \"report\": \"Report\",\n        \"acl\": \"ACL\",\n        \"mcl\": \"MCL\",\n        \"medial meniscus\": \"Medial Meniscus\",\n        \"lateral meniscus\": \"Lateral Meniscus\",\n        \"medial oa\": \"Medial OA\",\n        \"lateral oa\": \"Lateral OA\",\n        \"pf oa\": \"PF OA\",\n        \"effusion\": \"Effusion\",\n        \"synovitis\": \"Synovitis\",\n        \"baker's\": \"Baker's\",\n        \"bakers\": \"Baker's\",\n        \"contusion\": \"Contusion\",\n        \"fracture\": \"Fracture\",\n    }\n\n    rename_map = {}\n    for col in df.columns:\n        key = re.sub(r\"\\s+\", \" \", col.lower()).strip()\n        if key in canonical and canonical[key] not in df.columns:\n            rename_map[col] = canonical[key]\n\n    return df.rename(columns=rename_map)\n\n\ntrain = normalize_column_names(train)\ntest = normalize_column_names(test)\ntrain_series = normalize_column_names(train_series)\ntest_series = normalize_column_names(test_series)\nsample_submission = normalize_column_names(sample_submission)\n\nrequired_train_columns = [\"StudyInstanceUID\", \"Report\"] + TARGETS\nmissing_train_columns = [col for col in required_train_columns if col not in train.columns]\nif missing_train_columns:\n    raise KeyError(\n        \"Missing required columns in train.csv: \"\n        + \", \".join(missing_train_columns)\n        + f\"\\nAvailable columns: {train.columns.tolist()}\"\n    )\n\nfor df in (train, test):\n    if \"PatientSex\" not in df.columns:\n        df[\"PatientSex\"] = \"\"\n\nfor col in TARGETS:\n    train[col] = pd.to_numeric(train[col], errors=\"coerce\")\n\nexact_mask = train[TARGETS].notna().values\nfully_labeled_mask = exact_mask.all(axis=1)\nany_labeled_mask = exact_mask.any(axis=1)\n\nprint(f\"Train studies: {len(train):,}\")\nprint(f\"Test studies:  {len(test):,}\")\nprint(f\"Fully labeled studies: {fully_labeled_mask.sum():,}\")\nprint(f\"Studies with at least one exact label: {any_labeled_mask.sum():,}\")\nprint(f\"Train series: {len(train_series):,}\")\nprint(f\"Test series:  {len(test_series):,}\")\nprint(\"Optional PatientSex present in source:\", train[\"PatientSex\"].astype(str).str.len().gt(0).any())\nprint(\"Train columns:\", train.columns.tolist())\n\npreview_columns = [\"StudyInstanceUID\", \"Report\"] + TARGETS\nif train[\"PatientSex\"].astype(str).str.len().gt(0).any():\n    preview_columns.insert(1, \"PatientSex\")\n\ndisplay(train.loc[fully_labeled_mask, preview_columns].head(3))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Expanded multilingual report supervision\n\nThe report teacher combines exact labels, conservative multilingual rules and high-confidence rule-derived training examples. Exact labels retain full weight; report-derived image labels are deliberately down-weighted.\n\nFor OOF branch selection, the report teacher is re-fitted without the held-out exact studies, preventing the most direct label leakage.","metadata":{}},{"cell_type":"code","source":"def normalize_text(text):\n    text = \"\" if pd.isna(text) else str(text)\n    text = unicodedata.normalize(\"NFKD\", text)\n    text = \"\".join(ch for ch in text if not unicodedata.combining(ch))\n    text = text.lower()\n    text = text.replace(\"’\", \"'\").replace(\"–\", \"-\").replace(\"—\", \"-\")\n    text = re.sub(r\"\\s+\", \" \", text)\n    return text.strip()\n\n\nNEGATION = re.compile(\n    r\"\\b(?:no|not|without|absent|absence of|negative for|free of|\"\n    r\"sin|ausencia de|no se observa|no se evidencia|\"\n    r\"sem|ausencia|nao ha|nao evidencia|\"\n    r\"kein|keine|keinen|ohne|nicht nachweisbar|\"\n    r\"pas de|sans|absence de|aucun|aucune|\"\n    r\"senza|assenza di|non si osserva|\"\n    r\"geen|zonder|afwezig)\\b\"\n)\n\nUNCERTAIN = re.compile(\n    r\"\\b(?:possible|possibly|probable|probably|suspect|suspected|suggestive|\"\n    r\"cannot exclude|may represent|questionable|equivocal|\"\n    r\"posible|probable|sugestivo|sospecha|no se descarta|\"\n    r\"possivel|provavel|suspeita|\"\n    r\"moglich|verdacht|\"\n    r\"possible|probable|suspecte|\"\n    r\"possibile|probabile|sospetto)\\b\"\n)\n\nNORMAL_WORDS = re.compile(\n    r\"\\b(?:intact|normal|preserved|unremarkable|stable|\"\n    r\"integro|integros|conservado|conservados|normal|\"\n    r\"intacto|intactos|preservado|\"\n    r\"intakt|unauffallig|erhalten|\"\n    r\"intacte|intacts|normal|\"\n    r\"integro|integri|conservato)\\b\"\n)\n\nINJURY_WORDS = re.compile(\n    r\"\\b(?:tear|torn|rupture|ruptured|sprain|injury|disruption|insufficiency|\"\n    r\"rotura|ruptura|desgarro|lesion|esguince|distension|\"\n    r\"rottura|lesione|\"\n    r\"dechirure|rupture|entorse|lesion|\"\n    r\"riss|ruptur|zerrung|verletzung|\"\n    r\"scheur|ruptuur)\\w*\"\n)\n\nMENISCUS_TEAR_WORDS = re.compile(\n    r\"\\b(?:tear|torn|rupture|ruptured|fissure|cleavage|radial|bucket handle|root tear|\"\n    r\"rotura|ruptura|desgarro|fisura|lesion|\"\n    r\"rottura|fissurazione|lesione|\"\n    r\"dechirure|fissure|rupture|lesion|\"\n    r\"riss|ruptur|einriss|\"\n    r\"scheur|ruptuur)\\w*\"\n)\n\nOA_WORDS = re.compile(\n    r\"\\b(?:osteoarthrit|osteoarthros|arthros|gonarthros|degenerative joint|\"\n    r\"cartilage loss|chondral loss|joint space narrowing|chondromalacia|\"\n    r\"artrosis|osteoartritis|gonartrosis|degeneracion condral|condropatia|\"\n    r\"osteoartrose|artrose|condropatia|\"\n    r\"arthrose|gonarthrose|knorpelschaden|\"\n    r\"arthrose|chondropathie|\"\n    r\"artrosi|gonartrosi|condropatia)\\w*\"\n)\n\nALIASES = {\n    \"ACL\": re.compile(\n        r\"\\b(?:acl|anterior cruciate ligament|ligamento cruzado anterior|\"\n        r\"ligamento crociato anteriore|ligament croise anterieur|\"\n        r\"vorderes kreuzband|voorste kruisband)\\b\"\n    ),\n    \"MCL\": re.compile(\n        r\"\\b(?:mcl|medial collateral ligament|ligamento colateral medial|\"\n        r\"ligamento collaterale mediale|ligament collateral medial|\"\n        r\"mediales kollateralband|mediale band)\\b\"\n    ),\n    \"Medial Meniscus\": re.compile(\n        r\"\\b(?:medial menisc\\w*|menisc\\w* medial\\w*|menisco medial|\"\n        r\"menisco interno|menisque medial|innenmeniskus|menisco mediale)\\b\"\n    ),\n    \"Lateral Meniscus\": re.compile(\n        r\"\\b(?:lateral menisc\\w*|menisc\\w* lateral\\w*|menisco lateral|\"\n        r\"menisco externo|menisque lateral|aussenmeniskus|menisco laterale)\\b\"\n    ),\n    \"Medial OA\": re.compile(\n        r\"\\b(?:medial compartment|compartimento medial|compartimento interno|\"\n        r\"femorotibial medial|mediales kompartiment|compartiment medial)\\b\"\n    ),\n    \"Lateral OA\": re.compile(\n        r\"\\b(?:lateral compartment|compartimento lateral|compartimento externo|\"\n        r\"femorotibial lateral|laterales kompartiment|compartiment lateral)\\b\"\n    ),\n    \"PF OA\": re.compile(\n        r\"\\b(?:patellofemoral|patello femoral|femoropatellar|femoro patelar|\"\n        r\"femororotulian|retropatellar|femoropatellaire|femoropatellare)\\w*\"\n    ),\n    \"Effusion\": re.compile(\n        r\"\\b(?:joint effusion|effusion|hydarthrosis|derrame articular|derrame|\"\n        r\"liquido articular|derrame articular|erguss|gelenkerguss|\"\n        r\"epanchement|versamento articolare)\\b\"\n    ),\n    \"Synovitis\": re.compile(\n        r\"\\b(?:synovitis|sinovitis|synovite|synovial inflammation|\"\n        r\"synoviale entzundung)\\b\"\n    ),\n    \"Baker's\": re.compile(\n        r\"\\b(?:baker'?s? cyst|popliteal cyst|quiste de baker|quiste popliteo|\"\n        r\"cisto de baker|kyste de baker|baker zyste|cisti di baker)\\b\"\n    ),\n    \"Contusion\": re.compile(\n        r\"\\b(?:bone contusion|bone bruise|marrow contusion|traumatic marrow edema|\"\n        r\"contusion osea|contusao ossea|edema ose\\w*|edema de medula|\"\n        r\"contusion osseuse|oedeme osseux|knochenmarkodem|knochenkontusion)\\b\"\n    ),\n    \"Fracture\": re.compile(\n        r\"\\b(?:fracture|fractura|fratura|fraktur|frattura)\\w*\"\n    ),\n}\n\nGROUP_NORMAL = {\n    \"ACL\": re.compile(r\"\\b(?:cruciate ligaments?|ligamentos cruzados|kreuzbander)\\b\"),\n    \"MCL\": re.compile(r\"\\b(?:collateral ligaments?|ligamentos colaterales|kollateralbander)\\b\"),\n    \"Medial Meniscus\": re.compile(r\"\\b(?:menisci|meniscos|menisken|menisques)\\b\"),\n    \"Lateral Meniscus\": re.compile(r\"\\b(?:menisci|meniscos|menisken|menisques)\\b\"),\n}\n\n\ndef split_report(text):\n    text = normalize_text(text)\n    parts = re.split(r\"(?<=[\\.\\!\\?;:])\\s+|\\n+\", text)\n    return [p.strip() for p in parts if p.strip()]\n\n\ndef local_window(sentence, start, end, radius=90):\n    lo = max(0, start - radius)\n    hi = min(len(sentence), end + radius)\n    return sentence[lo:hi]\n\n\ndef target_evidence(sentence, target):\n    alias = ALIASES[target]\n    evidences = []\n\n    for m in alias.finditer(sentence):\n        window = local_window(sentence, m.start(), m.end())\n        negated = bool(NEGATION.search(window))\n        uncertain = bool(UNCERTAIN.search(window))\n        normal = bool(NORMAL_WORDS.search(window))\n\n        if target in (\"ACL\", \"MCL\"):\n            pathology = bool(INJURY_WORDS.search(window))\n        elif target in (\"Medial Meniscus\", \"Lateral Meniscus\"):\n            pathology = bool(MENISCUS_TEAR_WORDS.search(window))\n        elif target in (\"Medial OA\", \"Lateral OA\", \"PF OA\"):\n            pathology = bool(OA_WORDS.search(window))\n        else:\n            pathology = True\n\n        if negated or (normal and not pathology):\n            evidences.append((0.03, 0.95))\n        elif pathology:\n            if uncertain:\n                evidences.append((0.68, 0.45))\n            else:\n                base_conf = 0.82 if target == \"Contusion\" else 0.95\n                evidences.append((0.97, base_conf))\n\n    if target in GROUP_NORMAL:\n        group = GROUP_NORMAL[target]\n        for m in group.finditer(sentence):\n            window = local_window(sentence, m.start(), m.end())\n            if NORMAL_WORDS.search(window) or NEGATION.search(window):\n                evidences.append((0.03, 0.85))\n\n    if target in (\"Medial OA\", \"Lateral OA\", \"PF OA\"):\n        tri = re.search(r\"\\b(?:tricompartmental|tri-compartmental|tricompartimental)\\b\", sentence)\n        if tri and OA_WORDS.search(sentence):\n            evidences.append((0.97, 0.90))\n\n    return evidences\n\n\ndef report_rule_scores(report):\n    sentences = split_report(report)\n    probs = np.full(len(TARGETS), 0.5, dtype=np.float32)\n    confs = np.zeros(len(TARGETS), dtype=np.float32)\n\n    for j, target in enumerate(TARGETS):\n        candidates = []\n        for i, sentence in enumerate(sentences):\n            for p, c in target_evidence(sentence, target):\n                # Findings near the end often occur in the impression/conclusion.\n                position_bonus = 0.05 * (i / max(1, len(sentences) - 1))\n                candidates.append((min(1.0, c + position_bonus), p))\n        if candidates:\n            candidates.sort(reverse=True)\n            confs[j], probs[j] = candidates[0]\n    return probs, confs\n\n\nreports = train[\"Report\"].fillna(\"\").astype(str).tolist()\nrule_probs = np.zeros((len(train), len(TARGETS)), dtype=np.float32)\nrule_conf = np.zeros_like(rule_probs)\n\nfor i, report in enumerate(tqdm(reports, desc=\"Rule pseudo-labels\")):\n    rule_probs[i], rule_conf[i] = report_rule_scores(report)\n\nrule_rows = []\nfor j, target in enumerate(TARGETS):\n    known = train[target].notna().values\n    covered = known & (rule_conf[:, j] >= 0.50)\n    if covered.sum() > 0:\n        pred = (rule_probs[covered, j] >= 0.5).astype(int)\n        truth = train.loc[covered, target].astype(int).values\n        acc = float((pred == truth).mean())\n    else:\n        acc = np.nan\n    rule_rows.append({\n        \"target\": target,\n        \"exact_n\": int(known.sum()),\n        \"rule_covered\": int(covered.sum()),\n        \"rule_accuracy_on_covered\": acc\n    })\n\ndisplay(pd.DataFrame(rule_rows))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"normalized_reports = [normalize_text(x) for x in reports]\n\ntext_vectorizer = FeatureUnion([\n    (\"char\", TfidfVectorizer(\n        analyzer=\"char_wb\", ngram_range=(3, 5), min_df=2,\n        max_features=60_000, sublinear_tf=True, dtype=np.float32\n    )),\n    (\"word\", TfidfVectorizer(\n        analyzer=\"word\", ngram_range=(1, 2), min_df=1,\n        max_features=20_000, sublinear_tf=True, dtype=np.float32\n    )),\n])\n\nX_text = text_vectorizer.fit_transform(normalized_reports)\nprint(\"TF-IDF shape:\", X_text.shape)\n\n\ndef fit_text_probabilities(exclude_indices=None, compute_oof=True):\n    exclude = np.zeros(len(train), dtype=bool)\n    if exclude_indices is not None:\n        exclude[np.asarray(exclude_indices, dtype=int)] = True\n\n    text_probs = np.full((len(train), len(TARGETS)), 0.5, dtype=np.float32)\n    text_oof = np.full_like(text_probs, np.nan)\n    text_auc = {target: np.nan for target in TARGETS}\n    c_values = (1.2, 4.0)\n\n    for j, target in enumerate(TARGETS):\n        known = train[target].notna().values\n        exact_idx = np.where(known & (~exclude))[0]\n        exact_y = train.loc[exact_idx, target].astype(int).values\n\n        if len(exact_idx) == 0:\n            continue\n        if len(np.unique(exact_y)) < 2:\n            text_probs[:, j] = float(exact_y.mean())\n            continue\n\n        rule_idx = np.where(\n            (~known)\n            & (~exclude)\n            & (rule_conf[:, j] >= 0.86)\n            & (np.abs(rule_probs[:, j] - 0.5) >= 0.42)\n        )[0]\n        rule_y = (rule_probs[rule_idx, j] >= 0.5).astype(int)\n\n        def fit_predict(train_exact_idx, pred_idx):\n            train_parts = [np.asarray(train_exact_idx, dtype=int)]\n            y_parts = [\n                train.loc[train_exact_idx, target].astype(int).values\n            ]\n            weight_parts = [\n                np.ones(len(train_exact_idx), dtype=np.float32)\n            ]\n\n            if len(rule_idx):\n                train_parts.append(rule_idx)\n                y_parts.append(rule_y)\n                weight_parts.append(\n                    np.full(\n                        len(rule_idx),\n                        CFG.rule_teacher_weight,\n                        dtype=np.float32\n                    )\n                )\n\n            fit_idx = np.concatenate(train_parts)\n            fit_y = np.concatenate(y_parts)\n            fit_weight = np.concatenate(weight_parts)\n\n            predictions = []\n            for c_value in c_values:\n                model = LogisticRegression(\n                    C=c_value,\n                    solver=\"liblinear\",\n                    class_weight=\"balanced\",\n                    max_iter=1400,\n                    random_state=CFG.seed + j + int(10 * c_value)\n                )\n                model.fit(\n                    X_text[fit_idx],\n                    fit_y,\n                    sample_weight=fit_weight\n                )\n                predictions.append(\n                    model.predict_proba(X_text[pred_idx])[:, 1]\n                )\n            return np.mean(predictions, axis=0)\n\n        if compute_oof:\n            minority = int(np.bincount(exact_y).min())\n            n_splits = min(3, minority)\n            if n_splits >= 2:\n                skf = StratifiedKFold(\n                    n_splits=n_splits,\n                    shuffle=True,\n                    random_state=CFG.seed + j\n                )\n                for tr_local, va_local in skf.split(exact_idx, exact_y):\n                    tr_idx = exact_idx[tr_local]\n                    va_idx = exact_idx[va_local]\n                    text_oof[va_idx, j] = fit_predict(tr_idx, va_idx)\n\n                valid = ~np.isnan(text_oof[exact_idx, j])\n                if valid.sum() and len(np.unique(exact_y[valid])) == 2:\n                    text_auc[target] = roc_auc_score(\n                        exact_y[valid],\n                        text_oof[exact_idx[valid], j]\n                    )\n\n        text_probs[:, j] = fit_predict(\n            exact_idx,\n            np.arange(len(train))\n        )\n\n    return text_probs, text_oof, text_auc\n\n\ndef estimate_rule_reliability():\n    quality = np.full(len(TARGETS), 0.55, dtype=np.float32)\n    stats = []\n\n    for j, target in enumerate(TARGETS):\n        known = exact_mask[:, j]\n        covered = known & (rule_conf[:, j] >= 0.50)\n        y = train.loc[covered, target].astype(int).values\n        p = (rule_probs[covered, j] >= 0.5).astype(int)\n\n        if len(y) >= 6 and len(np.unique(y)) == 2:\n            bacc = balanced_accuracy_score(y, p)\n            reliability = np.clip((bacc - 0.50) / 0.35, 0.10, 1.00)\n        elif len(y) > 0:\n            bacc = float((y == p).mean())\n            reliability = np.clip((bacc - 0.50) / 0.40, 0.15, 0.80)\n        else:\n            bacc = np.nan\n            reliability = 0.35\n\n        quality[j] = reliability\n        stats.append({\n            \"target\": target,\n            \"rule_eval_n\": int(len(y)),\n            \"rule_balanced_accuracy\": bacc,\n            \"rule_reliability\": reliability\n        })\n\n    return quality, pd.DataFrame(stats)\n\n\ndef make_supervision(exclude_indices=None, compute_oof=True, text_reliability_override=None):\n    text_probs, text_oof, text_auc = fit_text_probabilities(\n        exclude_indices=exclude_indices,\n        compute_oof=compute_oof\n    )\n    rule_reliability, rule_stats = estimate_rule_reliability()\n\n    if text_reliability_override is None:\n        text_reliability = np.asarray([\n            np.clip((text_auc.get(target, np.nan) - 0.50) / 0.30, 0.0, 1.0)\n            if np.isfinite(text_auc.get(target, np.nan)) else 0.0\n            for target in TARGETS\n        ], dtype=np.float32)\n    else:\n        text_reliability = np.asarray(\n            text_reliability_override, dtype=np.float32\n        ).copy()\n\n    rule_strength = rule_conf * rule_reliability[None, :]\n    text_extremeness = np.clip(2.0 * np.abs(text_probs - 0.5), 0.0, 1.0)\n    text_strength = text_extremeness * text_reliability[None, :]\n\n    denom = rule_strength + 0.65 * text_strength\n    blended = np.full_like(text_probs, 0.5)\n\n    valid = denom > 1e-6\n    blended[valid] = (\n        rule_strength[valid] * rule_probs[valid]\n        + 0.65 * text_strength[valid] * text_probs[valid]\n    ) / denom[valid]\n\n    overall_reliability = np.maximum(rule_strength, text_strength)\n    pseudo_probs = 0.5 + overall_reliability * (blended - 0.5)\n    distance = np.abs(pseudo_probs - 0.5)\n\n    rule_selected = (\n        (rule_strength >= 0.42)\n        & (np.abs(rule_probs - 0.5) >= CFG.rule_threshold)\n    )\n    text_selected = (\n        (text_reliability[None, :] >= 0.25)\n        & (np.abs(text_probs - 0.5) >= CFG.text_threshold)\n    )\n    pseudo_mask = rule_selected | text_selected\n\n    pseudo_weight = (\n        CFG.min_pseudo_weight\n        + 0.60 * overall_reliability\n        + 0.25 * (2.0 * distance)\n    ).clip(CFG.min_pseudo_weight, CFG.max_pseudo_weight).astype(np.float32)\n\n    targets = pseudo_probs.astype(np.float32)\n    mask = pseudo_mask.astype(np.float32)\n    weight = (\n        CFG.pseudo_global_scale * pseudo_weight * mask\n    ).astype(np.float32)\n\n    exact_values = train[TARGETS].fillna(0.5).values.astype(np.float32)\n    targets[exact_mask] = exact_values[exact_mask]\n    mask[exact_mask] = 1.0\n    weight[exact_mask] = 1.0\n\n    diagnostics = rule_stats.copy()\n    diagnostics[\"text_oof_auc\"] = [\n        text_auc.get(target, np.nan) for target in TARGETS\n    ]\n    diagnostics[\"text_reliability\"] = text_reliability\n    return targets, mask, weight, text_oof, text_auc, diagnostics\n\n\n(\n    targets_all,\n    target_mask_all,\n    target_weight_all,\n    text_oof,\n    text_auc,\n    supervision_diagnostics\n) = make_supervision(\n    exclude_indices=None,\n    compute_oof=CFG.compute_text_oof\n)\n\ndisplay(supervision_diagnostics)\n\nGLOBAL_TEXT_RELIABILITY = np.asarray([\n    np.clip((text_auc.get(target, np.nan) - 0.50) / 0.30, 0.0, 1.0)\n    if np.isfinite(text_auc.get(target, np.nan)) else 0.0\n    for target in TARGETS\n], dtype=np.float32)\n\ncoverage = target_mask_all.mean(axis=0)\nsummary = pd.DataFrame({\n    \"target\": TARGETS,\n    \"exact_count\": exact_mask.sum(axis=0),\n    \"training_coverage\": coverage,\n    \"soft_positive_rate\": [\n        np.average(\n            targets_all[:, j],\n            weights=np.maximum(target_weight_all[:, j], 1e-6)\n        )\n        for j in range(len(TARGETS))\n    ]\n})\ndisplay(summary)","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fully_idx = np.where(fully_labeled_mask)[0]\nexact_indices = np.where(exact_mask.any(axis=1))[0]\n\nprint(\"Fully labeled studies:\", len(fully_idx))\nprint(\"Studies with at least one exact label:\", len(exact_indices))\n\n\ndef balanced_multilabel_folds(indices, y, n_splits, seed):\n    indices = np.asarray(indices, dtype=int)\n    y = np.asarray(y, dtype=np.float32)\n\n    if len(indices) == 0:\n        return []\n\n    n_splits = max(1, min(int(n_splits), len(indices)))\n    if n_splits == 1:\n        return [indices.copy()]\n\n    rng = np.random.default_rng(seed)\n    y = np.nan_to_num(y, nan=0.0)\n    y = (y > 0.5).astype(np.float32)\n\n    n_samples, n_targets = y.shape\n    target_sizes = np.full(\n        n_splits,\n        n_samples // n_splits,\n        dtype=np.int64\n    )\n    target_sizes[: n_samples % n_splits] += 1\n\n    positive_total = y.sum(axis=0)\n    rarity = 1.0 / np.maximum(positive_total, 1.0)\n    global_prevalence = positive_total / max(n_samples, 1)\n\n    priority = (\n        (y * rarity[None, :]).sum(axis=1)\n        + 0.03 * y.sum(axis=1)\n        + 1e-4 * rng.random(n_samples)\n    )\n    order = np.argsort(-priority)\n\n    assignments = np.full(n_samples, -1, dtype=np.int64)\n    fold_size = np.zeros(n_splits, dtype=np.int64)\n    fold_positive = np.zeros(\n        (n_splits, n_targets),\n        dtype=np.float64\n    )\n\n    # Seed every fold with one high-priority sample.\n    # This makes an empty validation fold impossible.\n    for fold, local_idx in enumerate(order[:n_splits]):\n        assignments[local_idx] = fold\n        fold_size[fold] += 1\n        fold_positive[fold] += y[local_idx]\n\n    desired_positive = (\n        positive_total[None, :]\n        * target_sizes[:, None]\n        / max(n_samples, 1)\n    )\n\n    for local_idx in order[n_splits:]:\n        sample = y[local_idx]\n        available = np.where(fold_size < target_sizes)[0]\n\n        if len(available) == 0:\n            available = np.arange(n_splits)\n\n        costs = []\n        for fold in available:\n            candidate_positive = fold_positive[fold] + sample\n            candidate_size = fold_size[fold] + 1\n\n            prevalence = (\n                candidate_positive\n                / max(candidate_size, 1)\n            )\n            prevalence_cost = np.sum(\n                rarity\n                * (prevalence - global_prevalence) ** 2\n            )\n\n            positive_cost = np.sum(\n                rarity\n                * (\n                    candidate_positive\n                    - desired_positive[fold]\n                ) ** 2\n                / (desired_positive[fold] + 1.0)\n            )\n\n            fill_ratio = (\n                candidate_size\n                / max(target_sizes[fold], 1)\n            )\n            size_cost = 0.08 * fill_ratio ** 2\n\n            costs.append(\n                0.65 * prevalence_cost\n                + 0.35 * positive_cost\n                + size_cost\n            )\n\n        costs = np.asarray(costs)\n        minimum = costs.min()\n        tied = available[\n            np.isclose(costs, minimum, rtol=1e-8, atol=1e-10)\n        ]\n        chosen = int(rng.choice(tied))\n\n        assignments[local_idx] = chosen\n        fold_size[chosen] += 1\n        fold_positive[chosen] += sample\n\n    folds = [\n        indices[assignments == fold]\n        for fold in range(n_splits)\n    ]\n\n    if any(len(fold) == 0 for fold in folds):\n        raise RuntimeError(\n            f\"Internal fold construction failure: \"\n            f\"{[len(fold) for fold in folds]}\"\n        )\n\n    if max(map(len, folds)) - min(map(len, folds)) > 1:\n        raise RuntimeError(\n            f\"Unbalanced folds: {[len(fold) for fold in folds]}\"\n        )\n\n    return folds\n\n\nif len(fully_idx) >= 2 and CFG.run_oof_blending:\n    requested_folds = min(\n        CFG.oof_folds,\n        len(fully_idx)\n    )\n    fully_targets = train.loc[\n        fully_idx, TARGETS\n    ].astype(int).values\n\n    OOF_FOLDS = balanced_multilabel_folds(\n        fully_idx,\n        fully_targets,\n        requested_folds,\n        CFG.seed\n    )\nelse:\n    OOF_FOLDS = [fully_idx] if len(fully_idx) else []\n\nOOF_FOLDS = [\n    np.asarray(fold, dtype=int)\n    for fold in OOF_FOLDS\n    if len(fold) > 0\n]\n\nprint(\"OOF fold sizes:\", [len(fold) for fold in OOF_FOLDS])\n\nif len(OOF_FOLDS) >= 2:\n    fold_rows = []\n    for fold_index, fold_indices in enumerate(OOF_FOLDS):\n        fold_y = train.loc[\n            fold_indices, TARGETS\n        ].astype(int).values\n        row = {\n            \"fold\": fold_index,\n            \"size\": len(fold_indices)\n        }\n        row.update({\n            target: int(fold_y[:, target_index].sum())\n            for target_index, target in enumerate(TARGETS)\n        })\n        fold_rows.append(row)\n\n    display(pd.DataFrame(fold_rows))\nelse:\n    print(\"OOF blending disabled: fewer than two non-empty folds.\")","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Dual-sequence wide-context MRI representation\n\nFor each anatomical plane, v6 selects two distinct acquisitions when available:\n\n- **Fluid**: prioritizes fluid-sensitive and fat-suppressed sequences;\n- **Structural**: prioritizes non-fluid and non-fat-suppressed anatomy-focused sequences.\n\nEach selected series is represented by three physically ordered positions spread around the central volume. The three positions form one RGB input, retaining broader slice coverage without three separate encoder passes.","metadata":{}},{"cell_type":"code","source":"try:\n    import cv2\nexcept ImportError:\n    cv2 = None\n\n\ndef numeric_flag(value):\n    try:\n        return float(value)\n    except Exception:\n        return 0.0\n\n\ndef choose_distinct_series(part, contrast, used_ids):\n    fluid = part[\"Fluid_Sensitive\"].fillna(0).astype(float)\n    fat = part[\"Fat_Suppression\"].fillna(0).astype(float)\n\n    if contrast == \"Fluid\":\n        score = 3.0 * fluid + 2.0 * fat\n    else:\n        score = 2.7 * (1.0 - fluid) + 1.3 * (1.0 - fat)\n\n    ordered = part.assign(_score=score).sort_values(\n        \"_score\", ascending=False\n    )\n\n    for _, row in ordered.iterrows():\n        series_id = str(row[\"SeriesInstanceUID\"])\n        if series_id not in used_ids:\n            return row\n\n    return None\n\n\ndef build_sequence_slots(df_series):\n    result = {}\n    study_meta = {}\n\n    for study_id, rows in tqdm(\n        df_series.groupby(\"StudyInstanceUID\", sort=False),\n        total=df_series[\"StudyInstanceUID\"].nunique(),\n        desc=\"Selecting dual-sequence slots\"\n    ):\n        study_id = str(study_id)\n        selected = {}\n        plane_values = rows[\"Anatomical_Plane\"].astype(str).str.lower()\n\n        meta_values = []\n        for plane in PLANES:\n            part = rows[plane_values == plane.lower()]\n            count = len(part)\n            fluid_count = (\n                part[\"Fluid_Sensitive\"].fillna(0).astype(float).sum()\n                if count else 0.0\n            )\n            fat_count = (\n                part[\"Fat_Suppression\"].fillna(0).astype(float).sum()\n                if count else 0.0\n            )\n            meta_values.extend([\n                np.log1p(count) / 3.0,\n                fluid_count / max(count, 1),\n                fat_count / max(count, 1),\n            ])\n\n            used_ids = set()\n            for contrast in CONTRASTS:\n                row = choose_distinct_series(part, contrast, used_ids)\n                slot_name = f\"{plane}_{contrast}\"\n                if row is None:\n                    continue\n\n                series_id = str(row[\"SeriesInstanceUID\"])\n                used_ids.add(series_id)\n                selected[slot_name] = {\n                    \"series_id\": series_id,\n                    \"plane\": plane,\n                    \"contrast\": contrast,\n                    \"fluid\": numeric_flag(row.get(\"Fluid_Sensitive\", 0)),\n                    \"fat\": numeric_flag(row.get(\"Fat_Suppression\", 0)),\n                }\n\n        total = len(rows)\n        meta_values.extend([\n            np.log1p(total) / 4.0,\n            rows[\"Fluid_Sensitive\"].fillna(0).astype(float).mean()\n            if total else 0.0,\n            rows[\"Fat_Suppression\"].fillna(0).astype(float).mean()\n            if total else 0.0,\n        ])\n\n        result[study_id] = selected\n        study_meta[study_id] = np.asarray(meta_values, dtype=np.float32)\n\n    return result, study_meta\n\n\nHEADER_TAGS = [\n    \"InstanceNumber\",\n    \"SliceLocation\",\n    \"ImagePositionPatient\",\n    \"ImageOrientationPatient\",\n]\n\n\ndef list_dicom_paths(folder):\n    try:\n        paths = [\n            Path(entry.path)\n            for entry in os.scandir(folder)\n            if entry.is_file() and entry.name.lower().endswith(\".dcm\")\n        ]\n    except Exception:\n        return []\n\n    paths.sort(key=lambda path: path.name)\n    return paths\n\n\ndef physical_coordinate(ds, fallback):\n    try:\n        orientation = np.asarray(\n            ds.ImageOrientationPatient, dtype=np.float64\n        )\n        position = np.asarray(\n            ds.ImagePositionPatient, dtype=np.float64\n        )\n        if orientation.size == 6 and position.size == 3:\n            normal = np.cross(orientation[:3], orientation[3:])\n            norm = np.linalg.norm(normal)\n            if norm > 1e-8:\n                return float(np.dot(position, normal / norm))\n    except Exception:\n        pass\n\n    for attribute in (\"SliceLocation\", \"InstanceNumber\"):\n        try:\n            value = float(getattr(ds, attribute))\n            if np.isfinite(value):\n                return value\n        except Exception:\n            pass\n\n    return float(fallback)\n\n\ndef sampled_wide_triplet(folder, plane, contrast):\n    paths = list_dicom_paths(folder)\n    n_files = len(paths)\n    if n_files == 0:\n        return [], 0\n\n    k = min(CFG.center_header_samples, n_files)\n    seed = zlib.crc32(str(folder).encode(\"utf-8\")) & 0xFFFFFFFF\n    rng = np.random.default_rng(seed)\n\n    if k == n_files:\n        candidate_indices = np.arange(n_files)\n    else:\n        candidate_indices = np.sort(\n            rng.choice(n_files, size=k, replace=False)\n        )\n\n    positioned = []\n    for fallback, path_idx in enumerate(candidate_indices):\n        path = paths[int(path_idx)]\n        try:\n            ds = pydicom.dcmread(\n                str(path),\n                stop_before_pixels=True,\n                specific_tags=HEADER_TAGS,\n                force=True\n            )\n            coordinate = physical_coordinate(ds, fallback)\n        except Exception:\n            coordinate = float(fallback)\n        positioned.append((coordinate, path))\n\n    positioned.sort(key=lambda item: item[0])\n    ordered = [item[1] for item in positioned]\n\n    if contrast == \"Structural\":\n        fractions = (0.34, 0.50, 0.66)\n    elif plane == \"Coronal\":\n        fractions = (0.28, 0.50, 0.72)\n    else:\n        fractions = (0.31, 0.50, 0.69)\n\n    selected = []\n    for fraction in fractions:\n        idx = int(round(fraction * (len(ordered) - 1)))\n        selected.append(ordered[np.clip(idx, 0, len(ordered) - 1)])\n\n    return selected, n_files\n\n\ndef read_dicom_slice(path):\n    try:\n        ds = pydicom.dcmread(str(path), force=True)\n        array = ds.pixel_array.astype(np.float32)\n    except Exception:\n        return None\n\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0) or 1.0)\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0) or 0.0)\n    array = array * slope + intercept\n\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        array = array.max() - array\n\n    finite = array[np.isfinite(array)]\n    if finite.size == 0:\n        return None\n\n    nonzero = finite[np.abs(finite) > 1e-8]\n    source = nonzero if nonzero.size >= 64 else finite\n    low, high = np.percentile(source, [1.0, 99.0])\n\n    if not np.isfinite(low) or not np.isfinite(high) or high <= low:\n        low, high = float(finite.min()), float(finite.max())\n\n    array = np.nan_to_num(\n        array, nan=low, posinf=high, neginf=low\n    )\n    array = np.clip(array, low, high)\n    array = (array - low) / max(high - low, 1e-6)\n    return array.astype(np.float32)\n\n\ndef crop_resize_plane(array, size):\n    mask = array > 0.03\n    height, width = array.shape\n\n    if mask.sum() > 0.02 * height * width:\n        ys, xs = np.where(mask)\n        y0, y1 = int(ys.min()), int(ys.max()) + 1\n        x0, x1 = int(xs.min()), int(xs.max()) + 1\n        margin_y = max(2, int(0.07 * (y1 - y0)))\n        margin_x = max(2, int(0.07 * (x1 - x0)))\n        array = array[\n            max(0, y0 - margin_y):min(height, y1 + margin_y),\n            max(0, x0 - margin_x):min(width, x1 + margin_x)\n        ]\n\n    height, width = array.shape\n    side = max(height, width)\n    padded = np.zeros((side, side), dtype=np.float32)\n    offset_y = (side - height) // 2\n    offset_x = (side - width) // 2\n    padded[\n        offset_y:offset_y + height,\n        offset_x:offset_x + width\n    ] = array\n\n    if cv2 is not None:\n        interpolation = (\n            cv2.INTER_AREA if side > size else cv2.INTER_LINEAR\n        )\n        resized = cv2.resize(\n            padded, (size, size), interpolation=interpolation\n        )\n    else:\n        tensor = torch.from_numpy(padded)[None, None]\n        resized = F.interpolate(\n            tensor,\n            size=(size, size),\n            mode=\"bilinear\",\n            align_corners=False\n        )[0, 0].numpy()\n\n    return np.clip(resized * 255.0, 0, 255).astype(np.uint8)\n\n\ndef make_wide_2p5d_image(paths):\n    channels = []\n\n    for path in paths:\n        array = read_dicom_slice(path)\n        if array is not None:\n            channels.append(\n                crop_resize_plane(array, CFG.img_size)\n            )\n\n    if not channels:\n        return None\n\n    while len(channels) < 3:\n        channels.append(channels[-1].copy())\n\n    return np.stack(channels[:3], axis=0)\n\n\nclass StudyDicomDataset(Dataset):\n    def __init__(\n        self,\n        studies,\n        sequence_slots,\n        study_meta,\n        series_root\n    ):\n        self.studies = studies.reset_index(drop=True)\n        self.sequence_slots = sequence_slots\n        self.study_meta = study_meta\n        self.series_root = Path(series_root)\n\n    def __len__(self):\n        return len(self.studies)\n\n    def __getitem__(self, idx):\n        row = self.studies.iloc[idx]\n        study_id = str(row[\"StudyInstanceUID\"])\n        selected = self.sequence_slots.get(study_id, {})\n\n        images = np.zeros(\n            (\n                CFG.n_slots, CFG.n_slices, 3,\n                CFG.img_size, CFG.img_size\n            ),\n            dtype=np.uint8\n        )\n        token_mask = np.zeros(\n            (CFG.n_slots, CFG.n_slices), dtype=np.bool_\n        )\n        slot_meta = np.zeros(\n            (CFG.n_slots, SLOT_META_DIM), dtype=np.float32\n        )\n\n        for slot_idx, slot_name in enumerate(SLOT_NAMES):\n            info = selected.get(slot_name)\n            if info is None:\n                continue\n\n            folder = (\n                self.series_root\n                / study_id\n                / str(info[\"series_id\"])\n            )\n            paths, n_files = sampled_wide_triplet(\n                folder,\n                info[\"plane\"],\n                info[\"contrast\"]\n            )\n            image = make_wide_2p5d_image(paths)\n\n            if image is None:\n                continue\n\n            images[slot_idx, 0] = image\n            token_mask[slot_idx, 0] = True\n            slot_meta[slot_idx] = np.asarray([\n                info[\"fluid\"],\n                info[\"fat\"],\n                1.0 if info[\"contrast\"] == \"Structural\" else 0.0,\n                np.log1p(n_files) / 6.0,\n            ], dtype=np.float32)\n\n        sex = str(row.get(\"PatientSex\", \"\")).strip().lower()\n        sex_id = (\n            1 if sex.startswith(\"m\")\n            else 2 if sex.startswith(\"f\")\n            else 0\n        )\n\n        global_meta = self.study_meta.get(\n            study_id,\n            np.zeros(STUDY_META_DIM, dtype=np.float32)\n        )\n\n        return (\n            torch.from_numpy(images),\n            torch.from_numpy(token_mask),\n            torch.from_numpy(slot_meta),\n            torch.from_numpy(global_meta.copy()),\n            torch.tensor(sex_id, dtype=torch.long),\n            study_id\n        )\n\n\ntrain_slots, train_study_meta = build_sequence_slots(train_series)\ntest_slots, test_study_meta = build_sequence_slots(test_series)\n\nslot_counts = {\n    slot: sum(slot in item for item in train_slots.values())\n    for slot in SLOT_NAMES\n}\ndisplay(pd.DataFrame({\n    \"slot\": list(slot_counts),\n    \"train_studies_with_slot\": list(slot_counts.values())\n}))","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import AutoModel\n\n\nclass DinoMultiLevelStatistics(nn.Module):\n    def __init__(self, backbone):\n        super().__init__()\n        self.backbone = backbone\n\n    @staticmethod\n    def pooled(tokens, include_max):\n        cls_token = tokens[:, 0]\n        patches = tokens[:, 1:]\n        values = [cls_token, patches.mean(dim=1)]\n        if include_max:\n            values.append(patches.amax(dim=1))\n        return values\n\n    def forward(self, pixel_values):\n        output = self.backbone(\n            pixel_values=pixel_values,\n            output_hidden_states=True\n        )\n        hidden_states = output.hidden_states\n        middle_index = max(1, len(hidden_states) // 2)\n        middle = hidden_states[middle_index]\n        final = hidden_states[-1]\n\n        statistics = []\n        statistics.extend(self.pooled(middle, include_max=False))\n        statistics.extend(self.pooled(final, include_max=True))\n        return torch.cat(statistics, dim=1)\n\n\ndef load_dinov2():\n    model_dir = find_dinov2_dir()\n    backbone = AutoModel.from_pretrained(\n        str(model_dir),\n        local_files_only=True,\n        trust_remote_code=False\n    )\n    hidden = int(backbone.config.hidden_size)\n\n    original_layers = None\n    if hasattr(backbone, \"encoder\") and hasattr(\n        backbone.encoder, \"layer\"\n    ):\n        original_layers = len(backbone.encoder.layer)\n        keep = min(CFG.dino_layers, original_layers)\n        backbone.encoder.layer = nn.ModuleList(\n            list(backbone.encoder.layer)[:keep]\n        )\n        backbone.config.num_hidden_layers = keep\n\n    for parameter in backbone.parameters():\n        parameter.requires_grad = False\n\n    feature_dim = 5 * hidden\n    model = DinoMultiLevelStatistics(backbone).eval().to(\n        CFG.device\n    )\n\n    if CFG.n_gpus > 1:\n        model = nn.DataParallel(\n            model,\n            device_ids=list(range(CFG.n_gpus))\n        )\n\n    print(\"DINOv2:\", model_dir)\n    print(\"Base hidden size:\", hidden)\n    print(\n        \"Output feature size:\",\n        feature_dim,\n        \"(mid CLS/mean + final CLS/mean/max)\"\n    )\n    print(\"DINO blocks:\", CFG.dino_layers, \"/\", original_layers)\n    print(\"Feature extraction GPUs:\", max(1, CFG.n_gpus))\n    return model, feature_dim\n\n\n@torch.inference_mode()\ndef encode_chunks(model, images):\n    mean = torch.tensor(\n        [0.485, 0.456, 0.406],\n        device=CFG.device\n    ).view(1, 3, 1, 1)\n    std = torch.tensor(\n        [0.229, 0.224, 0.225],\n        device=CFG.device\n    ).view(1, 3, 1, 1)\n\n    outputs = []\n\n    for start in range(\n        0, len(images), CFG.feature_batch_size\n    ):\n        batch = images[\n            start:start + CFG.feature_batch_size\n        ].to(\n            CFG.device, non_blocking=True\n        ).float().div_(255.0)\n        batch = (batch - mean) / std\n\n        with torch.autocast(\n            device_type=\"cuda\",\n            dtype=torch.float16,\n            enabled=CFG.amp\n        ):\n            features = model(batch)\n\n        outputs.append(features.float().cpu())\n\n    return torch.cat(outputs, dim=0)\n\n\ndef extract_features(\n    studies,\n    slots,\n    study_meta_map,\n    series_root,\n    prefix,\n    model,\n    feature_dim\n):\n    stem = f\"{prefix}_{CFG.feature_cache_tag}\"\n    feat_path = WORK / f\"{stem}_features.npy\"\n    mask_path = WORK / f\"{stem}_token_mask.npy\"\n    slot_meta_path = WORK / f\"{stem}_slot_meta.npy\"\n    study_meta_path = WORK / f\"{stem}_study_meta.npy\"\n    sex_path = WORK / f\"{stem}_sex.npy\"\n    ids_path = WORK / f\"{stem}_ids.npy\"\n    done_path = WORK / f\"{stem}.done\"\n\n    expected_shape = (\n        len(studies),\n        CFG.n_slots,\n        CFG.n_slices,\n        feature_dim\n    )\n    expected_mask_shape = expected_shape[:-1]\n\n    cache_files = (\n        feat_path, mask_path, slot_meta_path,\n        study_meta_path, sex_path, ids_path\n    )\n    cache_ok = (\n        all(path.exists() for path in cache_files)\n        and done_path.exists()\n    )\n\n    if cache_ok and not CFG.rebuild_features:\n        try:\n            cached = np.load(feat_path, mmap_mode=\"r\")\n            cached_mask = np.load(mask_path, mmap_mode=\"r\")\n            cached_slot_meta = np.load(\n                slot_meta_path, mmap_mode=\"r\"\n            )\n            cached_study_meta = np.load(\n                study_meta_path, mmap_mode=\"r\"\n            )\n            cached_sex = np.load(sex_path, mmap_mode=\"r\")\n            cached_ids = np.load(\n                ids_path, allow_pickle=True\n            )\n\n            valid_cache = (\n                cached.shape == expected_shape\n                and cached_mask.shape == expected_mask_shape\n                and cached_slot_meta.shape\n                == (len(studies), CFG.n_slots, SLOT_META_DIM)\n                and cached_study_meta.shape\n                == (len(studies), STUDY_META_DIM)\n                and cached_sex.shape == (len(studies),)\n                and len(cached_ids) == len(studies)\n            )\n            if valid_cache:\n                print(\n                    f\"Using completed cached {prefix} features:\",\n                    cached.shape\n                )\n                return (\n                    feat_path, mask_path, slot_meta_path,\n                    study_meta_path, sex_path, ids_path\n                )\n        except Exception as exception:\n            print(\n                f\"Ignoring invalid {prefix} cache:\",\n                repr(exception)\n            )\n\n    done_path.unlink(missing_ok=True)\n\n    dataset = StudyDicomDataset(\n        studies,\n        slots,\n        study_meta_map,\n        series_root\n    )\n    workers = max(\n        1,\n        min(\n            CFG.feature_workers,\n            os.cpu_count() or CFG.feature_workers\n        )\n    )\n    loader_batch_size = CFG.feature_loader_batch_size\n\n    print(\n        f\"{prefix}: {workers} DICOM threads, \"\n        f\"{CFG.n_slots} dual-sequence images per study\"\n    )\n\n    features = np.lib.format.open_memmap(\n        feat_path,\n        mode=\"w+\",\n        dtype=np.float16,\n        shape=expected_shape\n    )\n    masks = np.lib.format.open_memmap(\n        mask_path,\n        mode=\"w+\",\n        dtype=np.bool_,\n        shape=expected_mask_shape\n    )\n    slot_meta_memmap = np.lib.format.open_memmap(\n        slot_meta_path,\n        mode=\"w+\",\n        dtype=np.float16,\n        shape=(len(studies), CFG.n_slots, SLOT_META_DIM)\n    )\n    study_meta_memmap = np.lib.format.open_memmap(\n        study_meta_path,\n        mode=\"w+\",\n        dtype=np.float16,\n        shape=(len(studies), STUDY_META_DIM)\n    )\n    sexes = np.lib.format.open_memmap(\n        sex_path,\n        mode=\"w+\",\n        dtype=np.int8,\n        shape=(len(studies),)\n    )\n    ids = np.empty(len(studies), dtype=object)\n\n    def safe_get(index):\n        try:\n            return dataset[index]\n        except Exception:\n            study_id = str(\n                studies.iloc[index][\"StudyInstanceUID\"]\n            )\n            return (\n                torch.zeros(\n                    CFG.n_slots,\n                    CFG.n_slices,\n                    3,\n                    CFG.img_size,\n                    CFG.img_size,\n                    dtype=torch.uint8\n                ),\n                torch.zeros(\n                    CFG.n_slots,\n                    CFG.n_slices,\n                    dtype=torch.bool\n                ),\n                torch.zeros(\n                    CFG.n_slots,\n                    SLOT_META_DIM,\n                    dtype=torch.float32\n                ),\n                torch.zeros(\n                    STUDY_META_DIM,\n                    dtype=torch.float32\n                ),\n                torch.tensor(0, dtype=torch.long),\n                study_id\n            )\n\n    cursor = 0\n    starts = range(0, len(dataset), loader_batch_size)\n\n    with ThreadPoolExecutor(max_workers=workers) as pool:\n        for start in tqdm(\n            starts,\n            total=math.ceil(\n                len(dataset) / loader_batch_size\n            ),\n            desc=f\"{prefix} dual-sequence features\"\n        ):\n            stop = min(\n                start + loader_batch_size,\n                len(dataset)\n            )\n            indices = list(range(start, stop))\n\n            if workers == 1:\n                items = [safe_get(index) for index in indices]\n            else:\n                items = list(pool.map(safe_get, indices))\n\n            images = torch.stack(\n                [item[0] for item in items], dim=0\n            )\n            token_mask = torch.stack(\n                [item[1] for item in items], dim=0\n            )\n            slot_meta = torch.stack(\n                [item[2] for item in items], dim=0\n            )\n            study_meta = torch.stack(\n                [item[3] for item in items], dim=0\n            )\n            sex_id = torch.stack(\n                [item[4] for item in items], dim=0\n            )\n            study_ids = [item[5] for item in items]\n\n            batch_size, slots_count, slices_count = (\n                token_mask.shape\n            )\n            valid_flat = token_mask.reshape(-1).bool()\n            batch_features = torch.zeros(\n                batch_size * slots_count * slices_count,\n                feature_dim,\n                dtype=torch.float32\n            )\n\n            if valid_flat.any():\n                flat_images = images.reshape(\n                    batch_size * slots_count * slices_count,\n                    3,\n                    CFG.img_size,\n                    CFG.img_size\n                )\n                encoded = encode_chunks(\n                    model,\n                    flat_images[valid_flat]\n                )\n                batch_features[valid_flat] = encoded\n\n            end = cursor + batch_size\n            features[cursor:end] = (\n                batch_features.reshape(\n                    batch_size,\n                    CFG.n_slots,\n                    CFG.n_slices,\n                    feature_dim\n                ).numpy().astype(np.float16)\n            )\n            masks[cursor:end] = token_mask.numpy()\n            slot_meta_memmap[cursor:end] = (\n                slot_meta.numpy().astype(np.float16)\n            )\n            study_meta_memmap[cursor:end] = (\n                study_meta.numpy().astype(np.float16)\n            )\n            sexes[cursor:end] = (\n                sex_id.numpy().astype(np.int8)\n            )\n            ids[cursor:end] = np.asarray(\n                [str(value) for value in study_ids],\n                dtype=object\n            )\n            cursor = end\n\n            if CFG.debug and cursor >= 32:\n                break\n\n    features.flush()\n    masks.flush()\n    slot_meta_memmap.flush()\n    study_meta_memmap.flush()\n    sexes.flush()\n    np.save(ids_path, ids)\n\n    del (\n        features,\n        masks,\n        slot_meta_memmap,\n        study_meta_memmap,\n        sexes\n    )\n\n    if cursor == len(studies):\n        done_path.write_text(\n            f\"rows={cursor}\\nshape={expected_shape}\\n\",\n            encoding=\"utf-8\"\n        )\n        print(\n            f\"Completed {prefix} feature cache:\",\n            expected_shape\n        )\n    else:\n        print(\n            f\"{prefix} cache is partial \"\n            f\"({cursor}/{len(studies)}) and will not be reused.\"\n        )\n\n    return (\n        feat_path, mask_path, slot_meta_path,\n        study_meta_path, sex_path, ids_path\n    )\n\n\ndinov2, feature_dim = load_dinov2()\n\ntrain_feature_files = extract_features(\n    train,\n    train_slots,\n    train_study_meta,\n    BASE / \"train_series\",\n    \"train\",\n    dinov2,\n    feature_dim\n)\ntest_feature_files = extract_features(\n    test,\n    test_slots,\n    test_study_meta,\n    BASE / \"test_series\",\n    \"test\",\n    dinov2,\n    feature_dim\n)\n\ndel dinov2\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Correlation-aware target-conditioned sequence head\n\nThe head receives six sequence tokens, their sequence descriptors and study-level protocol metadata. During training, random available modalities are dropped and features receive small Gaussian noise.\n\nEach diagnosis attends to different sequences. A second Transformer across the twelve diagnosis tokens models co-occurrence such as OA with effusion/synovitis or ACL injury with contusion.","metadata":{}},{"cell_type":"code","source":"class FeatureDataset(Dataset):\n    def __init__(\n        self,\n        feature_files,\n        indices,\n        labels=None,\n        label_mask=None,\n        label_weight=None,\n        exact_cells=None,\n        train_mode=False\n    ):\n        (\n            feat_path,\n            mask_path,\n            slot_meta_path,\n            study_meta_path,\n            sex_path,\n            ids_path\n        ) = feature_files\n\n        self.features = np.load(feat_path, mmap_mode=\"r\")\n        self.token_mask = np.load(mask_path, mmap_mode=\"r\")\n        self.slot_meta = np.load(\n            slot_meta_path, mmap_mode=\"r\"\n        )\n        self.study_meta = np.load(\n            study_meta_path, mmap_mode=\"r\"\n        )\n        self.sex = np.load(sex_path, mmap_mode=\"r\")\n        self.ids = np.load(ids_path, allow_pickle=True)\n        self.indices = np.asarray(indices, dtype=np.int64)\n        self.labels = labels\n        self.label_mask = label_mask\n        self.label_weight = label_weight\n        self.exact_cells = exact_cells\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, item):\n        index = int(self.indices[item])\n\n        features = torch.from_numpy(\n            np.asarray(\n                self.features[index],\n                dtype=np.float32\n            ).copy()\n        )\n        token_mask = torch.from_numpy(\n            np.asarray(\n                self.token_mask[index],\n                dtype=np.bool_\n            ).copy()\n        )\n        slot_meta = torch.from_numpy(\n            np.asarray(\n                self.slot_meta[index],\n                dtype=np.float32\n            ).copy()\n        )\n        study_meta = torch.from_numpy(\n            np.asarray(\n                self.study_meta[index],\n                dtype=np.float32\n            ).copy()\n        )\n        sex = torch.tensor(\n            int(self.sex[index]), dtype=torch.long\n        )\n\n        common = (\n            features,\n            token_mask,\n            slot_meta,\n            study_meta,\n            sex\n        )\n\n        if self.train_mode:\n            target = torch.from_numpy(\n                self.labels[index].astype(np.float32)\n            )\n            mask = torch.from_numpy(\n                self.label_mask[index].astype(np.float32)\n            )\n            weight = torch.from_numpy(\n                self.label_weight[index].astype(np.float32)\n            )\n            exact = torch.from_numpy(\n                self.exact_cells[index].astype(np.bool_)\n            )\n            return (\n                *common,\n                target,\n                mask,\n                weight,\n                exact,\n                index\n            )\n\n        return (\n            *common,\n            str(self.ids[index]),\n            index\n        )\n\n\nPLANE_PRIOR = torch.tensor([\n    [1.20, 0.90, 0.70, 0.45, 0.05, 0.05],\n    [0.20, 0.10, 1.30, 1.00, 0.05, 0.05],\n    [0.95, 0.80, 1.10, 0.85, 0.15, 0.15],\n    [0.80, 0.70, 0.95, 0.80, 0.35, 0.30],\n    [0.40, 0.85, 0.75, 1.20, 0.15, 0.25],\n    [0.40, 0.85, 0.75, 1.20, 0.15, 0.25],\n    [0.35, 0.55, 0.20, 0.30, 1.00, 1.25],\n    [0.90, 0.30, 0.45, 0.20, 1.15, 0.45],\n    [0.85, 0.30, 0.50, 0.20, 1.10, 0.45],\n    [1.15, 0.60, 0.40, 0.20, 0.60, 0.30],\n    [1.10, 0.35, 1.05, 0.35, 1.00, 0.30],\n    [0.70, 0.80, 0.80, 0.90, 0.70, 0.80],\n], dtype=torch.float32)\n\n\nclass DualSequenceTargetHead(nn.Module):\n    def __init__(self, input_dim, n_targets=len(TARGETS)):\n        super().__init__()\n        hidden = CFG.hidden_dim\n\n        self.input_norm = nn.LayerNorm(input_dim)\n        self.input_projection = nn.Sequential(\n            nn.Linear(input_dim, hidden),\n            nn.GELU(),\n            nn.Dropout(CFG.dropout)\n        )\n        self.slot_embedding = nn.Embedding(\n            CFG.n_slots, hidden\n        )\n        self.slot_meta_projection = nn.Sequential(\n            nn.LayerNorm(SLOT_META_DIM),\n            nn.Linear(SLOT_META_DIM, hidden),\n            nn.GELU()\n        )\n        self.study_meta_projection = nn.Sequential(\n            nn.LayerNorm(STUDY_META_DIM),\n            nn.Linear(STUDY_META_DIM, hidden),\n            nn.GELU(),\n            nn.Linear(hidden, hidden)\n        )\n        self.sex_embedding = nn.Embedding(3, hidden)\n\n        sequence_layer = nn.TransformerEncoderLayer(\n            d_model=hidden,\n            nhead=5,\n            dim_feedforward=3 * hidden,\n            dropout=CFG.dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True\n        )\n        self.sequence_encoder = nn.TransformerEncoder(\n            sequence_layer,\n            num_layers=2\n        )\n\n        self.target_queries = nn.Parameter(\n            torch.randn(n_targets, hidden) * 0.02\n        )\n        self.sequence_bias = nn.Parameter(\n            0.32 * PLANE_PRIOR.clone()\n        )\n\n        target_layer = nn.TransformerEncoderLayer(\n            d_model=hidden,\n            nhead=5,\n            dim_feedforward=2 * hidden,\n            dropout=CFG.dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=True\n        )\n        self.target_encoder = nn.TransformerEncoder(\n            target_layer,\n            num_layers=1\n        )\n\n        self.target_mlp = nn.Sequential(\n            nn.LayerNorm(hidden),\n            nn.Linear(hidden, hidden),\n            nn.GELU(),\n            nn.Dropout(CFG.dropout),\n            nn.Linear(hidden, hidden),\n            nn.GELU(),\n            nn.Dropout(CFG.dropout)\n        )\n        self.output_weight = nn.Parameter(\n            torch.randn(n_targets, hidden)\n            / math.sqrt(hidden)\n        )\n        self.output_bias = nn.Parameter(\n            torch.zeros(n_targets)\n        )\n        self.direct_head = nn.Linear(\n            hidden, n_targets\n        )\n\n    def stochastic_mask(self, mask):\n        if (\n            not self.training\n            or CFG.modality_dropout <= 0\n        ):\n            return mask\n\n        keep = (\n            torch.rand(\n                mask.shape,\n                device=mask.device\n            ) > CFG.modality_dropout\n        )\n        dropped = mask & keep\n        empty = ~dropped.any(dim=1)\n\n        if empty.any():\n            for batch_index in torch.where(empty)[0]:\n                valid_indices = torch.where(\n                    mask[batch_index]\n                )[0]\n                if len(valid_indices):\n                    chosen = valid_indices[\n                        torch.randint(\n                            len(valid_indices),\n                            (1,),\n                            device=mask.device\n                        )\n                    ]\n                    dropped[batch_index, chosen] = True\n\n        return dropped\n\n    def forward(\n        self,\n        features,\n        token_mask,\n        slot_meta,\n        study_meta,\n        sex\n    ):\n        x = features[:, :, 0]\n        mask = token_mask[:, :, 0].bool()\n\n        safe_mask = mask.clone()\n        empty = ~safe_mask.any(dim=1)\n        if empty.any():\n            safe_mask[empty, 0] = True\n            x = x.clone()\n            x[empty, 0] = 0.0\n\n        attention_mask = self.stochastic_mask(\n            safe_mask\n        )\n\n        if self.training and CFG.feature_noise > 0:\n            x = x + CFG.feature_noise * torch.randn_like(x)\n\n        x = self.input_projection(\n            self.input_norm(x)\n        )\n        slot_ids = torch.arange(\n            CFG.n_slots,\n            device=x.device\n        ).unsqueeze(0)\n\n        x = (\n            x\n            + self.slot_embedding(slot_ids)\n            + 0.35 * self.slot_meta_projection(slot_meta)\n        )\n\n        x = self.sequence_encoder(\n            x,\n            src_key_padding_mask=~attention_mask\n        )\n\n        attention_logits = torch.einsum(\n            \"th,bsh->bts\",\n            self.target_queries,\n            x\n        ) / math.sqrt(x.shape[-1])\n        attention_logits = (\n            attention_logits\n            + self.sequence_bias.unsqueeze(0)\n        )\n        attention_logits = attention_logits.masked_fill(\n            ~attention_mask[:, None, :],\n            -1e4\n        )\n        attention = attention_logits.softmax(dim=-1)\n\n        target_features = torch.einsum(\n            \"bts,bsh->bth\",\n            attention,\n            x\n        )\n\n        valid_float = attention_mask.float().unsqueeze(-1)\n        global_feature = (\n            (x * valid_float).sum(dim=1)\n            / valid_float.sum(dim=1).clamp_min(1.0)\n        )\n        protocol_feature = self.study_meta_projection(\n            study_meta\n        )\n        patient_feature = (\n            protocol_feature\n            + self.sex_embedding(sex)\n        )\n\n        target_features = (\n            target_features\n            + 0.30 * global_feature[:, None, :]\n            + 0.25 * patient_feature[:, None, :]\n            + self.target_queries[None, :, :]\n        )\n        target_features = self.target_encoder(\n            target_features\n        )\n        target_features = self.target_mlp(\n            target_features\n        )\n\n        logits = (\n            target_features\n            * self.output_weight[None, :, :]\n        ).sum(dim=-1) / math.sqrt(\n            target_features.shape[-1]\n        )\n        logits = logits + self.output_bias\n\n        direct_logits = self.direct_head(\n            global_feature + 0.30 * patient_feature\n        )\n        return logits + 0.22 * direct_logits\n\n\ndef compute_pos_weight(\n    indices,\n    targets,\n    label_mask,\n    label_weight\n):\n    y = targets[indices]\n    mask = label_mask[indices]\n    weight = label_weight[indices]\n\n    positive = (y * mask * weight).sum(axis=0)\n    negative = (\n        (1.0 - y) * mask * weight\n    ).sum(axis=0)\n    ratio = np.sqrt(\n        (negative + 1.0) / (positive + 1.0)\n    )\n\n    return torch.tensor(\n        np.clip(ratio, 0.70, 4.50),\n        dtype=torch.float32,\n        device=CFG.device\n    )\n\n\ndef masked_bce(\n    logits,\n    targets,\n    mask,\n    weight,\n    pos_weight\n):\n    loss = F.binary_cross_entropy_with_logits(\n        logits,\n        targets,\n        reduction=\"none\",\n        pos_weight=pos_weight\n    )\n    effective = mask * weight\n    return (\n        (loss * effective).sum()\n        / effective.sum().clamp_min(1.0)\n    )\n\n\ndef pairwise_ranking_loss(\n    logits,\n    targets,\n    exact_cells\n):\n    losses = []\n    hard_targets = targets > 0.5\n\n    for target_index in range(logits.shape[1]):\n        valid = exact_cells[:, target_index]\n        positive = logits[\n            valid & hard_targets[:, target_index],\n            target_index\n        ]\n        negative = logits[\n            valid & (~hard_targets[:, target_index]),\n            target_index\n        ]\n\n        if len(positive) and len(negative):\n            differences = (\n                positive[:, None]\n                - negative[None, :]\n            )\n            losses.append(\n                F.softplus(-differences).mean()\n            )\n\n    if not losses:\n        return logits.new_tensor(0.0)\n\n    return torch.stack(losses).mean()\n\n\n@torch.no_grad()\ndef update_ema(ema_model, model, decay):\n    ema_parameters = dict(\n        ema_model.named_parameters()\n    )\n    model_parameters = dict(\n        model.named_parameters()\n    )\n\n    for name, parameter in model_parameters.items():\n        ema_parameters[name].mul_(decay).add_(\n            parameter.detach(),\n            alpha=1.0 - decay\n        )\n\n    ema_buffers = dict(ema_model.named_buffers())\n    model_buffers = dict(model.named_buffers())\n    for name, buffer in model_buffers.items():\n        ema_buffers[name].copy_(buffer)\n\n\nprint(\"Study feature dimension:\", feature_dim)","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_indices = np.arange(len(train))\nexact_targets = train[TARGETS].fillna(0.5).values.astype(\n    np.float32\n)\nexact_training_mask = exact_mask.astype(np.float32)\nexact_training_weight = exact_mask.astype(np.float32)\n\n\ndef make_loader(\n    indices,\n    labels,\n    masks,\n    weights,\n    exact_cells,\n    seed,\n    batch_size,\n    steps,\n    boost_exact=False\n):\n    indices = np.asarray(indices, dtype=int)\n    dataset = FeatureDataset(\n        train_feature_files,\n        indices,\n        labels=labels,\n        label_mask=masks,\n        label_weight=weights,\n        exact_cells=exact_cells,\n        train_mode=True\n    )\n\n    sample_weights = np.ones(\n        len(indices), dtype=np.float64\n    )\n    if boost_exact:\n        sample_weights[\n            exact_mask[indices].any(axis=1)\n        ] = CFG.exact_sample_boost\n\n    sampler = WeightedRandomSampler(\n        sample_weights,\n        num_samples=batch_size * steps,\n        replacement=True,\n        generator=torch.Generator().manual_seed(seed)\n    )\n\n    return DataLoader(\n        dataset,\n        batch_size=batch_size,\n        sampler=sampler,\n        num_workers=0,\n        pin_memory=True,\n        drop_last=True\n    )\n\n\ndef train_epochs(\n    model,\n    ema_model,\n    loader,\n    optimizer,\n    epochs,\n    pos_weight,\n    rank_weight,\n    ema_decay,\n    seed,\n    label\n):\n    scaler = torch.cuda.amp.GradScaler(\n        enabled=CFG.amp\n    )\n    total_steps = max(1, epochs * len(loader))\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=total_steps,\n        eta_min=optimizer.param_groups[0][\"lr\"] * 0.05\n    )\n\n    for epoch in range(1, epochs + 1):\n        model.train()\n        running = 0.0\n\n        for step, batch in enumerate(loader):\n            (\n                features,\n                token_mask,\n                slot_meta,\n                study_meta,\n                sex,\n                target,\n                target_mask,\n                target_weight,\n                exact_cells_batch,\n                _\n            ) = batch\n\n            features = features.to(\n                CFG.device, non_blocking=True\n            )\n            token_mask = token_mask.to(\n                CFG.device, non_blocking=True\n            )\n            slot_meta = slot_meta.to(\n                CFG.device, non_blocking=True\n            )\n            study_meta = study_meta.to(\n                CFG.device, non_blocking=True\n            )\n            sex = sex.to(\n                CFG.device, non_blocking=True\n            )\n            target = target.to(\n                CFG.device, non_blocking=True\n            )\n            target_mask = target_mask.to(\n                CFG.device, non_blocking=True\n            )\n            target_weight = target_weight.to(\n                CFG.device, non_blocking=True\n            )\n            exact_cells_batch = exact_cells_batch.to(\n                CFG.device, non_blocking=True\n            )\n\n            optimizer.zero_grad(set_to_none=True)\n\n            with torch.autocast(\n                \"cuda\",\n                dtype=torch.float16,\n                enabled=CFG.amp\n            ):\n                logits = model(\n                    features,\n                    token_mask,\n                    slot_meta,\n                    study_meta,\n                    sex\n                )\n                bce = masked_bce(\n                    logits,\n                    target,\n                    target_mask,\n                    target_weight,\n                    pos_weight\n                )\n                ranking = pairwise_ranking_loss(\n                    logits,\n                    target,\n                    exact_cells_batch\n                )\n                loss = bce + rank_weight * ranking\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(\n                model.parameters(), 0.9\n            )\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n            update_ema(\n                ema_model,\n                model,\n                ema_decay\n            )\n            running += float(loss.item())\n\n            if CFG.debug and step >= 2:\n                break\n\n        print(\n            f\"{label} epoch {epoch:02d}/{epochs} \"\n            f\"loss={running / max(1, step + 1):.4f}\"\n        )\n\n    del scaler, scheduler\n\n\ndef train_weak_model(\n    seed,\n    training_indices,\n    supervision,\n    save_path=None\n):\n    labels, masks, weights = supervision\n\n    train_loader = make_loader(\n        training_indices,\n        labels,\n        masks,\n        weights,\n        exact_mask,\n        seed,\n        CFG.train_batch_size,\n        CFG.steps_per_epoch,\n        boost_exact=True\n    )\n\n    exact_train_indices = np.asarray([\n        index\n        for index in exact_indices\n        if index in set(np.asarray(training_indices, dtype=int))\n    ], dtype=int)\n\n    exact_batch_size = min(\n        128,\n        max(32, len(exact_train_indices))\n    )\n    exact_loader = make_loader(\n        exact_train_indices,\n        exact_targets,\n        exact_training_mask,\n        exact_training_weight,\n        exact_mask,\n        seed + 10_000,\n        exact_batch_size,\n        CFG.finetune_steps,\n        boost_exact=False\n    )\n\n    seed_everything(seed)\n    model = DualSequenceTargetHead(\n        feature_dim\n    ).to(CFG.device)\n    ema_model = deepcopy(model).eval()\n    for parameter in ema_model.parameters():\n        parameter.requires_grad = False\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.lr,\n        weight_decay=CFG.weight_decay\n    )\n    pos_weight = compute_pos_weight(\n        training_indices,\n        labels,\n        masks,\n        weights\n    )\n    train_epochs(\n        model,\n        ema_model,\n        train_loader,\n        optimizer,\n        CFG.pretrain_epochs,\n        pos_weight,\n        CFG.rank_loss_weight,\n        CFG.ema_decay,\n        seed,\n        f\"weak seed={seed}\"\n    )\n\n    model.load_state_dict(\n        ema_model.state_dict()\n    )\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.finetune_lr,\n        weight_decay=CFG.weight_decay\n    )\n    pos_weight_exact = compute_pos_weight(\n        exact_train_indices,\n        exact_targets,\n        exact_training_mask,\n        exact_training_weight\n    )\n    train_epochs(\n        model,\n        ema_model,\n        exact_loader,\n        optimizer,\n        CFG.finetune_epochs,\n        pos_weight_exact,\n        CFG.finetune_rank_weight,\n        0.990,\n        seed,\n        f\"exact-correction seed={seed}\"\n    )\n\n    if save_path is not None:\n        state = {\n            name: value.detach().cpu().clone()\n            for name, value in ema_model.state_dict().items()\n        }\n        torch.save({\n            \"state_dict\": state,\n            \"feature_dim\": feature_dim,\n            \"targets\": TARGETS,\n            \"seed\": seed,\n            \"branch\": \"weak\"\n        }, save_path)\n\n    del model, optimizer\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    return ema_model\n\n\ndef train_exact_model(\n    seed,\n    training_indices,\n    save_path=None\n):\n    training_indices = np.asarray(\n        training_indices, dtype=int\n    )\n    batch_size = min(\n        128,\n        max(32, len(training_indices))\n    )\n    loader = make_loader(\n        training_indices,\n        exact_targets,\n        exact_training_mask,\n        exact_training_weight,\n        exact_mask,\n        seed,\n        batch_size,\n        CFG.exact_only_steps,\n        boost_exact=False\n    )\n\n    seed_everything(seed)\n    model = DualSequenceTargetHead(\n        feature_dim\n    ).to(CFG.device)\n    ema_model = deepcopy(model).eval()\n    for parameter in ema_model.parameters():\n        parameter.requires_grad = False\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.exact_only_lr,\n        weight_decay=CFG.weight_decay\n    )\n    pos_weight = compute_pos_weight(\n        training_indices,\n        exact_targets,\n        exact_training_mask,\n        exact_training_weight\n    )\n    train_epochs(\n        model,\n        ema_model,\n        loader,\n        optimizer,\n        CFG.exact_only_epochs,\n        pos_weight,\n        CFG.exact_only_rank_weight,\n        0.992,\n        seed,\n        f\"exact-only seed={seed}\"\n    )\n\n    if save_path is not None:\n        state = {\n            name: value.detach().cpu().clone()\n            for name, value in ema_model.state_dict().items()\n        }\n        torch.save({\n            \"state_dict\": state,\n            \"feature_dim\": feature_dim,\n            \"targets\": TARGETS,\n            \"seed\": seed,\n            \"branch\": \"exact\"\n        }, save_path)\n\n    del model, optimizer\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    return ema_model\n\n\n@torch.no_grad()\ndef predict_indices(model, indices):\n    indices = np.asarray(indices, dtype=int)\n\n    if len(indices) == 0:\n        return (\n            np.empty(\n                (0, len(TARGETS)),\n                dtype=np.float32\n            ),\n            indices\n        )\n\n    dataset = FeatureDataset(\n        train_feature_files,\n        indices,\n        train_mode=False\n    )\n    loader = DataLoader(\n        dataset,\n        batch_size=CFG.train_batch_size,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=True\n    )\n\n    model.eval()\n    outputs = []\n    output_indices = []\n\n    for batch in loader:\n        (\n            features,\n            token_mask,\n            slot_meta,\n            study_meta,\n            sex,\n            _,\n            batch_indices\n        ) = batch\n\n        features = features.to(\n            CFG.device, non_blocking=True\n        )\n        token_mask = token_mask.to(\n            CFG.device, non_blocking=True\n        )\n        slot_meta = slot_meta.to(\n            CFG.device, non_blocking=True\n        )\n        study_meta = study_meta.to(\n            CFG.device, non_blocking=True\n        )\n        sex = sex.to(\n            CFG.device, non_blocking=True\n        )\n\n        with torch.autocast(\n            \"cuda\",\n            dtype=torch.float16,\n            enabled=CFG.amp\n        ):\n            logits = model(\n                features,\n                token_mask,\n                slot_meta,\n                study_meta,\n                sex\n            )\n\n        outputs.append(\n            logits.float().cpu().numpy()\n        )\n        output_indices.extend(\n            batch_indices.numpy().tolist()\n        )\n\n    if not outputs:\n        return (\n            np.empty(\n                (0, len(TARGETS)),\n                dtype=np.float32\n            ),\n            np.asarray(output_indices, dtype=int)\n        )\n\n    return (\n        np.concatenate(outputs, axis=0),\n        np.asarray(output_indices, dtype=int)\n    )\n\n\noof_weak = np.full(\n    (len(train), len(TARGETS)),\n    np.nan,\n    dtype=np.float32\n)\noof_exact = np.full_like(oof_weak, np.nan)\n\nif CFG.run_oof_blending and len(OOF_FOLDS) >= 2:\n    for fold_index, validation_indices in enumerate(\n        OOF_FOLDS\n    ):\n        validation_indices = np.asarray(\n            validation_indices,\n            dtype=int\n        )\n\n        if len(validation_indices) == 0:\n            print(\n                f\"Skipping empty OOF fold {fold_index + 1}.\"\n            )\n            continue\n\n        print(\n            f\"\\n===== OOF fold {fold_index + 1}/\"\n            f\"{len(OOF_FOLDS)} \"\n            f\"| validation={len(validation_indices)} =====\"\n        )\n        validation_set = set(validation_indices.tolist())\n        training_indices = np.asarray([\n            index\n            for index in all_indices\n            if index not in validation_set\n        ], dtype=int)\n        exact_fold_train = np.asarray([\n            index\n            for index in exact_indices\n            if index not in validation_set\n        ], dtype=int)\n\n        if len(training_indices) == 0:\n            print(\n                f\"Skipping fold {fold_index + 1}: \"\n                \"empty training set.\"\n            )\n            continue\n\n        if len(exact_fold_train) == 0:\n            print(\n                f\"Skipping fold {fold_index + 1}: \"\n                \"empty exact-label training set.\"\n            )\n            continue\n\n        fold_supervision_full = make_supervision(\n            exclude_indices=validation_indices,\n            compute_oof=False,\n            text_reliability_override=GLOBAL_TEXT_RELIABILITY\n        )\n        fold_supervision = (\n            fold_supervision_full[0],\n            fold_supervision_full[1],\n            fold_supervision_full[2]\n        )\n\n        weak_model = train_weak_model(\n            CFG.seed + 100 * fold_index,\n            training_indices,\n            fold_supervision\n        )\n        weak_logits, predicted_indices = predict_indices(\n            weak_model,\n            validation_indices\n        )\n        oof_weak[predicted_indices] = weak_logits\n        del weak_model\n\n        exact_model = train_exact_model(\n            CFG.seed + 100 * fold_index + 37,\n            exact_fold_train\n        )\n        exact_logits, predicted_indices = predict_indices(\n            exact_model,\n            validation_indices\n        )\n        oof_exact[predicted_indices] = exact_logits\n        del exact_model\n\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n\ndef oof_auc_table():\n    rows = []\n    blend_weights = np.full(\n        len(TARGETS), 0.70, dtype=np.float32\n    )\n    grid = np.linspace(0.0, 1.0, 11)\n\n    for target_index, target in enumerate(TARGETS):\n        valid = (\n            fully_labeled_mask\n            & np.isfinite(oof_weak[:, target_index])\n            & np.isfinite(oof_exact[:, target_index])\n        )\n        y = train.loc[\n            valid, target\n        ].astype(int).values\n\n        weak_auc = np.nan\n        exact_auc = np.nan\n        best_auc = np.nan\n        best_weight = 0.70\n\n        if valid.sum() and len(np.unique(y)) == 2:\n            weak_values = oof_weak[\n                valid, target_index\n            ]\n            exact_values = oof_exact[\n                valid, target_index\n            ]\n            weak_auc = roc_auc_score(\n                y, weak_values\n            )\n            exact_auc = roc_auc_score(\n                y, exact_values\n            )\n\n            weak_rank = rankdata(\n                weak_values, method=\"average\"\n            )\n            exact_rank = rankdata(\n                exact_values, method=\"average\"\n            )\n\n            for weak_weight in grid:\n                blended = (\n                    weak_weight * weak_rank\n                    + (1.0 - weak_weight) * exact_rank\n                )\n                auc = roc_auc_score(y, blended)\n                if (\n                    not np.isfinite(best_auc)\n                    or auc > best_auc\n                ):\n                    best_auc = auc\n                    best_weight = float(weak_weight)\n\n            best_weight = (\n                0.68 * best_weight\n                + 0.32 * 0.70\n            )\n\n        blend_weights[target_index] = best_weight\n        rows.append({\n            \"target\": target,\n            \"weak_oof_auc\": weak_auc,\n            \"exact_oof_auc\": exact_auc,\n            \"best_blend_oof_auc\": best_auc,\n            \"final_weak_weight\": best_weight\n        })\n\n    return pd.DataFrame(rows), blend_weights\n\n\nblend_table, TARGET_WEAK_WEIGHTS = oof_auc_table()\ndisplay(blend_table)\n\nvalid_macro = blend_table[\n    \"best_blend_oof_auc\"\n].dropna()\nif len(valid_macro):\n    print(\n        \"OOF macro AUC:\",\n        float(valid_macro.mean())\n    )\nelse:\n    print(\n        \"OOF unavailable; using default weak weight 0.70.\"\n    )\n\n\nfinal_supervision = (\n    targets_all,\n    target_mask_all,\n    target_weight_all\n)\n\nweak_model_paths = []\nfor seed in CFG.final_weak_seeds:\n    path = WORK / f\"knee_v6_weak_seed{seed}.pt\"\n    model = train_weak_model(\n        seed,\n        all_indices,\n        final_supervision,\n        save_path=path\n    )\n    del model\n    weak_model_paths.append(path)\n\nexact_model_paths = []\nfor seed in CFG.final_exact_seeds:\n    path = WORK / f\"knee_v6_exact_seed{seed}.pt\"\n    model = train_exact_model(\n        seed,\n        exact_indices,\n        save_path=path\n    )\n    del model\n    exact_model_paths.append(path)\n\nprint(\"Final weak models:\", weak_model_paths)\nprint(\"Final exact models:\", exact_model_paths)\nprint(\"Target weak weights:\", TARGET_WEAK_WEIGHTS)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef predict_test_model(model, loader):\n    model.eval()\n    logits_all = []\n    ids_all = []\n\n    for batch in tqdm(loader, leave=False):\n        (\n            features,\n            token_mask,\n            slot_meta,\n            study_meta,\n            sex,\n            study_ids,\n            _\n        ) = batch\n\n        features = features.to(\n            CFG.device, non_blocking=True\n        )\n        token_mask = token_mask.to(\n            CFG.device, non_blocking=True\n        )\n        slot_meta = slot_meta.to(\n            CFG.device, non_blocking=True\n        )\n        study_meta = study_meta.to(\n            CFG.device, non_blocking=True\n        )\n        sex = sex.to(\n            CFG.device, non_blocking=True\n        )\n\n        with torch.autocast(\n            \"cuda\",\n            dtype=torch.float16,\n            enabled=CFG.amp\n        ):\n            logits = model(\n                features,\n                token_mask,\n                slot_meta,\n                study_meta,\n                sex\n            )\n\n        logits_all.append(\n            logits.float().cpu().numpy()\n        )\n        ids_all.extend(list(study_ids))\n\n    return np.concatenate(\n        logits_all, axis=0\n    ), ids_all\n\n\ntest_dataset = FeatureDataset(\n    test_feature_files,\n    np.arange(len(test)),\n    train_mode=False\n)\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=CFG.train_batch_size,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=True\n)\n\n\ndef predict_branch(model_paths):\n    branch_logits = []\n    prediction_ids = None\n\n    for model_path in model_paths:\n        checkpoint = torch.load(\n            model_path, map_location=\"cpu\"\n        )\n        model = DualSequenceTargetHead(\n            checkpoint[\"feature_dim\"]\n        ).to(CFG.device)\n        model.load_state_dict(\n            checkpoint[\"state_dict\"]\n        )\n\n        logits, ids = predict_test_model(\n            model, test_loader\n        )\n        branch_logits.append(logits)\n\n        if prediction_ids is None:\n            prediction_ids = ids\n\n        del model\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    return np.stack(\n        branch_logits, axis=0\n    ), prediction_ids\n\n\nweak_logits, prediction_ids = predict_branch(\n    weak_model_paths\n)\nexact_logits, exact_ids = predict_branch(\n    exact_model_paths\n)\n\nassert prediction_ids == exact_ids\nn_test = len(test)\n\n\ndef branch_rank_ensemble(logits_stack):\n    probability_mean = (\n        1.0 / (1.0 + np.exp(-logits_stack))\n    ).mean(axis=0)\n    rank_mean = np.zeros_like(\n        probability_mean\n    )\n\n    for target_index in range(len(TARGETS)):\n        ranks = []\n        for model_index in range(\n            len(logits_stack)\n        ):\n            rank = rankdata(\n                logits_stack[\n                    model_index,\n                    :,\n                    target_index\n                ],\n                method=\"average\"\n            )\n            ranks.append(\n                rank / (n_test + 1.0)\n            )\n        rank_mean[:, target_index] = np.mean(\n            ranks, axis=0\n        )\n\n    return (\n        0.88 * rank_mean\n        + 0.12 * probability_mean\n    )\n\n\nweak_predictions = branch_rank_ensemble(\n    weak_logits\n)\nexact_predictions = branch_rank_ensemble(\n    exact_logits\n)\n\ntest_probs = (\n    TARGET_WEAK_WEIGHTS[None, :]\n    * weak_predictions\n    + (\n        1.0 - TARGET_WEAK_WEIGHTS[None, :]\n    )\n    * exact_predictions\n)\n\nprediction_frame = pd.DataFrame(\n    test_probs,\n    columns=TARGETS\n)\nprediction_frame.insert(\n    0,\n    \"StudyInstanceUID\",\n    prediction_ids\n)\n\nsubmission = sample_submission[\n    [\"StudyInstanceUID\"]\n].merge(\n    prediction_frame,\n    on=\"StudyInstanceUID\",\n    how=\"left\"\n)\n\nfor target in TARGETS:\n    submission[target] = (\n        submission[target]\n        .fillna(0.5)\n        .clip(1e-5, 1.0 - 1e-5)\n    )\n\nsubmission.to_csv(\n    WORK / \"submission.csv\",\n    index=False\n)\nsubmission.to_csv(\n    \"submission.csv\",\n    index=False\n)\n\nprint(\"submission.csv:\", submission.shape)\nprint(\n    \"Weak models:\", len(weak_logits),\n    \"| Exact models:\", len(exact_logits)\n)\ndisplay(\n    pd.DataFrame({\n        \"target\": TARGETS,\n        \"weak_weight\": TARGET_WEAK_WEIGHTS\n    })\n)\ndisplay(submission.head())","metadata":{"lines_to_next_cell":2},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Why v6 should improve over the 0.621 model\n\nThe 0.621 notebook used one fluid-prioritized series per plane and a single final-layer DINO representation. v6 adds complementary structural sequences, broader within-series coverage, intermediate/final DINO statistics, protocol metadata, modality dropout and explicit label-correlation modeling.\n\nThe OOF stage does not assume that report-derived supervision helps every diagnosis equally. For every target it compares:\n\n1. a large weakly supervised image model;\n2. a small exact-label-only image model;\n3. rank blends of both branches.\n\nThe selected blend is shrunk toward the more stable weak branch to reduce overfitting to the small labeled set.\n\nExpected runtime is materially higher than v5 because six 126×126 images and eight DINO blocks are processed per study. The default profile is still safely within the nine-hour notebook limit on two T4 GPUs, but actual runtime depends on DICOM transfer syntax and Kaggle storage throughput.\n\n## v6.1 fold stability fix\n\nThe earlier custom greedy splitter could assign zero studies to one OOF fold.\nThat produced an empty `DataLoader`, followed by\n`np.concatenate([])` in `predict_indices`.\n\nv6.1 guarantees:\n\n- every OOF fold is non-empty;\n- fold sizes differ by at most one;\n- empty prediction requests return an empty matrix safely;\n- invalid training/validation folds are skipped with a clear message;\n- the v6 DINO feature-cache tag is unchanged, so completed features are reused.\n","metadata":{}}]}