{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# RSNA KNEE ABNORMALITY DETECTION\n# PHASE 1: DATASET INTELLIGENCE / DATA AUDIT\n# ============================================================\n\nimport os\nimport re\nimport gc\nimport json\nimport time\nimport math\nimport random\nimport hashlib\nimport warnings\nfrom pathlib import Path\nfrom collections import Counter, defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\n# ------------------------------------------------------------\n# Audit configuration\n# ------------------------------------------------------------\n\nN_STUDY_AUDIT = 60\nN_DICOM_AUDIT = 80\nN_VISUAL_SERIES = 12\nN_VISUAL_SLICES = 9\n\n# Do not increase these blindly.\n# Phase 1 is deliberately sampling-based.\n\nOUTPUT_DIR = Path(\"/kaggle/working/rsna_knee_audit\")\nVISUAL_DIR = OUTPUT_DIR / \"visual_audit\"\n\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\nVISUAL_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"Output directory:\", OUTPUT_DIR)\nprint(\"Seed:\", SEED)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:37:50.700344Z","iopub.execute_input":"2026-08-11T07:37:50.701682Z","iopub.status.idle":"2026-08-11T07:37:50.712182Z","shell.execute_reply.started":"2026-08-11T07:37:50.701627Z","shell.execute_reply":"2026-08-11T07:37:50.710964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# KAGGLE ENVIRONMENT CHECK\n# ============================================================\n\nprint(\"Current working directory:\")\nprint(Path.cwd())\n\nprint(\"\\nKaggle input exists:\", Path(\"/kaggle/input\").exists())\nprint(\"Kaggle working exists:\", Path(\"/kaggle/working\").exists())\n\nif Path(\"/kaggle/input\").exists():\n    input_dirs = [p for p in Path(\"/kaggle/input\").iterdir() if p.is_dir()]\n    \n    print(\"\\nDatasets mounted under /kaggle/input:\")\n    for p in input_dirs:\n        print(\" -\", p)\n\nprint(\"\\nPython version:\")\nimport sys\nprint(sys.version)\n\nprint(\"\\nNumPy:\", np.__version__)\nprint(\"Pandas:\", pd.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:37:50.714097Z","iopub.execute_input":"2026-08-11T07:37:50.714543Z","iopub.status.idle":"2026-08-11T07:37:50.743642Z","shell.execute_reply.started":"2026-08-11T07:37:50.714512Z","shell.execute_reply":"2026-08-11T07:37:50.742443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DATASET PATH DISCOVERY\n# ============================================================\n\nKAGGLE_INPUT = Path(\"/kaggle/input\")\n\nEXPECTED_FILES = [\n    \"train.csv\",\n    \"train_series.csv\",\n    \"test.csv\",\n    \"test_series.csv\",\n    \"sample_submission.csv\",\n]\n\ndef find_file(filename):\n    matches = list(KAGGLE_INPUT.rglob(filename))\n    \n    if not matches:\n        return None\n    \n    # Prefer paths where several expected competition files\n    # are located nearby.\n    scored = []\n    \n    for path in matches:\n        parent = path.parent\n        \n        score = 0\n        \n        for expected in EXPECTED_FILES:\n            if (parent / expected).exists():\n                score += 1\n        \n        if (parent / \"train_series\").exists():\n            score += 3\n        \n        if (parent / \"test_series\").exists():\n            score += 3\n        \n        scored.append((score, path))\n    \n    scored.sort(key=lambda x: x[0], reverse=True)\n    return scored[0][1]\n\n\nDATA_PATHS = {}\n\nfor filename in EXPECTED_FILES:\n    path = find_file(filename)\n    DATA_PATHS[filename] = path\n\nprint(\"Discovered competition files:\\n\")\n\nfor name, path in DATA_PATHS.items():\n    print(f\"{name:25s} -> {path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:37:50.778997Z","iopub.execute_input":"2026-08-11T07:37:50.779583Z","iopub.status.idle":"2026-08-11T07:51:06.60133Z","shell.execute_reply.started":"2026-08-11T07:37:50.779543Z","shell.execute_reply":"2026-08-11T07:51:06.6Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# DISCOVER TRAIN/TEST SERIES DIRECTORIES\n# ============================================================\n\ndef find_directory(dirname):\n    matches = [p for p in KAGGLE_INPUT.rglob(dirname) if p.is_dir()]\n    \n    if not matches:\n        return None\n    \n    # Prefer the directory closest to the discovered CSV files.\n    train_csv = DATA_PATHS.get(\"train.csv\")\n    \n    if train_csv:\n        train_parent = train_csv.parent\n        \n        same_parent = [\n            p for p in matches\n            if p.parent == train_parent\n        ]\n        \n        if same_parent:\n            return same_parent[0]\n    \n    return matches[0]\n\n\nTRAIN_SERIES_DIR = find_directory(\"train_series\")\nTEST_SERIES_DIR = find_directory(\"test_series\")\n\nprint(\"TRAIN_SERIES_DIR:\")\nprint(TRAIN_SERIES_DIR)\n\nprint(\"\\nTEST_SERIES_DIR:\")\nprint(TEST_SERIES_DIR)\n\nassert TRAIN_SERIES_DIR is not None, \"train_series directory was not found.\"\nassert TEST_SERIES_DIR is not None, \"test_series directory was not found.\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:51:06.604281Z","iopub.execute_input":"2026-08-11T07:51:06.604714Z","iopub.status.idle":"2026-08-11T07:54:32.79123Z","shell.execute_reply.started":"2026-08-11T07:51:06.604682Z","shell.execute_reply":"2026-08-11T07:54:32.790122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LOAD CSV FILES\n# ============================================================\n\ndef load_csv(name):\n    path = DATA_PATHS[name]\n    \n    if path is None:\n        raise FileNotFoundError(f\"{name} was not found.\")\n    \n    print(f\"Loading {name} from:\")\n    print(path)\n    \n    df = pd.read_csv(path)\n    \n    print(\"Shape:\", df.shape)\n    \n    return df\n\n\ntrain = load_csv(\"train.csv\")\ntrain_series = load_csv(\"train_series.csv\")\ntest = load_csv(\"test.csv\")\ntest_series = load_csv(\"test_series.csv\")\nsample_submission = load_csv(\"sample_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:32.792394Z","iopub.execute_input":"2026-08-11T07:54:32.792675Z","iopub.status.idle":"2026-08-11T07:54:33.028552Z","shell.execute_reply.started":"2026-08-11T07:54:32.792645Z","shell.execute_reply":"2026-08-11T07:54:33.027258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# BASIC DATASET STRUCTURE\n# ============================================================\n\nprint(\"TRAIN\")\nprint(train.head())\nprint(\"\\nColumns:\")\nprint(train.columns.tolist())\n\nprint(\"\\nTRAIN SERIES\")\nprint(train_series.head())\nprint(\"\\nColumns:\")\nprint(train_series.columns.tolist())\n\nprint(\"\\nTEST\")\nprint(test.head())\nprint(\"\\nColumns:\")\nprint(test.columns.tolist())\n\nprint(\"\\nTEST SERIES\")\nprint(test_series.head())\n\nprint(\"\\nSAMPLE SUBMISSION\")\nprint(sample_submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.032069Z","iopub.execute_input":"2026-08-11T07:54:33.032983Z","iopub.status.idle":"2026-08-11T07:54:33.064258Z","shell.execute_reply.started":"2026-08-11T07:54:33.032939Z","shell.execute_reply":"2026-08-11T07:54:33.063062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TARGET DISCOVERY\n# ============================================================\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n]\n\nmissing_targets = [t for t in TARGETS if t not in train.columns]\n\nif missing_targets:\n    print(\"WARNING: Expected target columns missing:\")\n    print(missing_targets)\n\nTARGETS_PRESENT = [t for t in TARGETS if t in train.columns]\n\nprint(\"\\nTargets found:\")\nfor i, target in enumerate(TARGETS_PRESENT, 1):\n    print(f\"{i:2d}. {target}\")\n\nassert len(TARGETS_PRESENT) == 12, (\n    f\"Expected 12 targets, found {len(TARGETS_PRESENT)}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.06679Z","iopub.execute_input":"2026-08-11T07:54:33.067615Z","iopub.status.idle":"2026-08-11T07:54:33.076623Z","shell.execute_reply.started":"2026-08-11T07:54:33.06758Z","shell.execute_reply":"2026-08-11T07:54:33.075693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# TRAIN CSV AUDIT\n# ============================================================\n\nsummary_rows = []\n\ndef add_summary(category, metric, value, notes=\"\"):\n    summary_rows.append({\n        \"category\": category,\n        \"metric\": metric,\n        \"value\": value,\n        \"notes\": notes\n    })\n\n\nadd_summary(\"train\", \"rows\", len(train))\nadd_summary(\"train\", \"columns\", len(train.columns))\nadd_summary(\n    \"train\",\n    \"unique_studies\",\n    train[\"StudyInstanceUID\"].nunique()\n)\nadd_summary(\n    \"train\",\n    \"duplicate_study_rows\",\n    int(train[\"StudyInstanceUID\"].duplicated().sum())\n)\n\nif \"PatientSex\" in train.columns:\n    add_summary(\n        \"train\",\n        \"missing_patient_sex\",\n        int(train[\"PatientSex\"].isna().sum())\n    )\n\nif \"Report\" in train.columns:\n    add_summary(\n        \"train\",\n        \"reports_present\",\n        int(train[\"Report\"].notna().sum())\n    )\n    \n    add_summary(\n        \"train\",\n        \"reports_missing\",\n        int(train[\"Report\"].isna().sum())\n    )\n\n\nprint(\"Train rows:\", len(train))\nprint(\"Unique studies:\", train[\"StudyInstanceUID\"].nunique())\nprint(\"Duplicate StudyInstanceUID rows:\",\n      train[\"StudyInstanceUID\"].duplicated().sum())\n\nprint(\"\\nMissing values:\")\nprint(train.isna().sum().sort_values(ascending=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.078006Z","iopub.execute_input":"2026-08-11T07:54:33.078389Z","iopub.status.idle":"2026-08-11T07:54:33.112273Z","shell.execute_reply.started":"2026-08-11T07:54:33.078362Z","shell.execute_reply":"2026-08-11T07:54:33.111265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# LABEL DISTRIBUTION\n# ============================================================\n\nlabel_rows = []\n\nfor target in TARGETS_PRESENT:\n    \n    series = pd.to_numeric(train[target], errors=\"coerce\")\n    \n    positive = int((series == 1).sum())\n    negative = int((series == 0).sum())\n    missing = int(series.isna().sum())\n    total_available = positive + negative\n    \n    prevalence = (\n        positive / total_available\n        if total_available > 0\n        else np.nan\n    )\n    \n    ratio = (\n        negative / positive\n        if positive > 0\n        else np.inf\n    )\n    \n    label_rows.append({\n        \"target\": target,\n        \"positive\": positive,\n        \"negative\": negative,\n        \"missing\": missing,\n        \"available_labels\": total_available,\n        \"prevalence\": prevalence,\n        \"negative_positive_ratio\": ratio\n    })\n    \n    add_summary(\n        \"labels\",\n        f\"{target}_positive\",\n        positive\n    )\n    \n    add_summary(\n        \"labels\",\n        f\"{target}_missing\",\n        missing\n    )\n\n\nlabel_distribution = pd.DataFrame(label_rows)\n\ndisplay(label_distribution)\n\nlabel_distribution.to_csv(\n    OUTPUT_DIR / \"label_distribution.csv\",\n    index=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.113629Z","iopub.execute_input":"2026-08-11T07:54:33.114104Z","iopub.status.idle":"2026-08-11T07:54:33.15051Z","shell.execute_reply.started":"2026-08-11T07:54:33.11407Z","shell.execute_reply":"2026-08-11T07:54:33.149277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 11: EXACT LABELED STUDY AUDIT\n# ============================================================\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n]\n\nlabel_mask = train[TARGETS].notna().all(axis=1)\n\nlabeled_train = train.loc[label_mask].copy()\n\nprint(\"=\" * 70)\nprint(\"EXPLICITLY LABELED STUDY AUDIT\")\nprint(\"=\" * 70)\n\nprint(\"\\nTotal train studies:\", len(train))\nprint(\"Fully labeled studies:\", len(labeled_train))\n\nprint(\"\\nLabel values:\")\nfor target in TARGETS:\n    print(\n        f\"{target:20s} \"\n        f\"unique={sorted(labeled_train[target].dropna().unique().tolist())}\"\n    )\n\nprint(\"\\nLabeled study IDs:\")\nprint(labeled_train[\"StudyInstanceUID\"].head(10).tolist())\n\n# Save the small labeled dataset.\nlabeled_train.to_csv(\n    OUTPUT_DIR / \"explicitly_labeled_studies.csv\",\n    index=False\n)\n\nprint(\n    \"\\nSaved:\",\n    OUTPUT_DIR / \"explicitly_labeled_studies.csv\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.152098Z","iopub.execute_input":"2026-08-11T07:54:33.152519Z","iopub.status.idle":"2026-08-11T07:54:33.17933Z","shell.execute_reply.started":"2026-08-11T07:54:33.152482Z","shell.execute_reply":"2026-08-11T07:54:33.177556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 12: LABELED TARGET RELATIONSHIPS\n# ============================================================\n\nlabeled_targets = labeled_train[TARGETS].astype(int)\n\nprint(\"=\" * 70)\nprint(\"TARGET CORRELATION\")\nprint(\"=\" * 70)\n\ncorr = labeled_targets.corr()\n\ndisplay(corr.round(3))\n\ncorr.to_csv(\n    OUTPUT_DIR / \"labeled_target_correlations.csv\"\n)\n\n# ------------------------------------------------------------\n# Positive abnormality combinations\n# ------------------------------------------------------------\n\nfrom collections import Counter\n\ncombination_counter = Counter()\n\nfor _, row in labeled_targets.iterrows():\n    \n    positive_targets = tuple(\n        target\n        for target in TARGETS\n        if row[target] == 1\n    )\n    \n    combination_counter[positive_targets] += 1\n\n\ncombination_rows = []\n\nfor combination, count in combination_counter.most_common():\n    \n    combination_rows.append({\n        \"combination\": \" | \".join(combination),\n        \"count\": count\n    })\n\n\ncombinations_df = pd.DataFrame(combination_rows)\n\nprint(\"\\nMost common abnormality combinations:\")\ndisplay(combinations_df.head(20))\n\ncombinations_df.to_csv(\n    OUTPUT_DIR / \"label_combinations.csv\",\n    index=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.182597Z","iopub.execute_input":"2026-08-11T07:54:33.18309Z","iopub.status.idle":"2026-08-11T07:54:33.245086Z","shell.execute_reply.started":"2026-08-11T07:54:33.183059Z","shell.execute_reply":"2026-08-11T07:54:33.243726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 13: INSPECT LABELED REPORTS\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"SAMPLE OF LABELED REPORTS\")\nprint(\"=\" * 70)\n\ndisplay_columns = [\n    \"StudyInstanceUID\",\n    \"Report\"\n] + TARGETS\n\n# Show all 58 because the dataset is tiny.\nfor i, (_, row) in enumerate(\n    labeled_train[display_columns].iterrows()\n):\n    \n    print(\"\\n\" + \"=\" * 100)\n    print(f\"LABELED STUDY {i + 1} / {len(labeled_train)}\")\n    print(\"=\" * 100)\n    \n    print(\"StudyInstanceUID:\")\n    print(row[\"StudyInstanceUID\"])\n    \n    print(\"\\nLabels:\")\n    \n    positive_labels = [\n        target\n        for target in TARGETS\n        if row[target] == 1\n    ]\n    \n    negative_labels = [\n        target\n        for target in TARGETS\n        if row[target] == 0\n    ]\n    \n    print(\"Positive:\", positive_labels)\n    print(\"Negative:\", negative_labels)\n    \n    print(\"\\nREPORT:\")\n    print(row[\"Report\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.248382Z","iopub.execute_input":"2026-08-11T07:54:33.24875Z","iopub.status.idle":"2026-08-11T07:54:33.284271Z","shell.execute_reply.started":"2026-08-11T07:54:33.24872Z","shell.execute_reply":"2026-08-11T07:54:33.282969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 14: REPORT SCRIPT / LANGUAGE AUDIT\n# ============================================================\n\ndef script_statistics(text):\n    \n    text = str(text)\n    \n    if not text.strip():\n        return {\n            \"ascii_letters\": 0,\n            \"latin_extended\": 0,\n            \"cyrillic\": 0,\n            \"greek\": 0,\n            \"other_letters\": 0,\n            \"total_letters\": 0\n        }\n    \n    ascii_letters = 0\n    latin_extended = 0\n    cyrillic = 0\n    greek = 0\n    other_letters = 0\n    \n    for char in text:\n        \n        code = ord(char)\n        \n        if (\"A\" <= char <= \"Z\") or (\"a\" <= char <= \"z\"):\n            ascii_letters += 1\n        \n        elif 0x00C0 <= code <= 0x024F:\n            latin_extended += 1\n        \n        elif 0x0400 <= code <= 0x04FF:\n            cyrillic += 1\n        \n        elif 0x0370 <= code <= 0x03FF:\n            greek += 1\n        \n        elif char.isalpha():\n            other_letters += 1\n    \n    \n    total_letters = (\n        ascii_letters\n        + latin_extended\n        + cyrillic\n        + greek\n        + other_letters\n    )\n    \n    return {\n        \"ascii_letters\": ascii_letters,\n        \"latin_extended\": latin_extended,\n        \"cyrillic\": cyrillic,\n        \"greek\": greek,\n        \"other_letters\": other_letters,\n        \"total_letters\": total_letters\n    }\n\n\nlanguage_rows = []\n\nfor _, row in train.iterrows():\n    \n    stats = script_statistics(row[\"Report\"])\n    \n    total = stats[\"total_letters\"]\n    \n    if total == 0:\n        dominant = \"UNKNOWN\"\n    \n    elif (\n        stats[\"cyrillic\"] / total > 0.5\n    ):\n        dominant = \"CYRILLIC\"\n    \n    elif (\n        stats[\"greek\"] / total > 0.5\n    ):\n        dominant = \"GREEK\"\n    \n    elif (\n        (\n            stats[\"ascii_letters\"]\n            + stats[\"latin_extended\"]\n        ) / total > 0.8\n    ):\n        dominant = \"LATIN\"\n    \n    else:\n        dominant = \"MIXED/OTHER\"\n    \n    \n    language_rows.append({\n        \"StudyInstanceUID\": row[\"StudyInstanceUID\"],\n        \"dominant_script\": dominant,\n        \"ascii_letters\": stats[\"ascii_letters\"],\n        \"latin_extended\": stats[\"latin_extended\"],\n        \"cyrillic\": stats[\"cyrillic\"],\n        \"greek\": stats[\"greek\"],\n        \"other_letters\": stats[\"other_letters\"],\n        \"total_letters\": total,\n        \"report_words\": len(str(row[\"Report\"]).split())\n    })\n\n\nreport_script_audit = pd.DataFrame(language_rows)\n\nprint(\"Dominant report script:\")\ndisplay(\n    report_script_audit[\"dominant_script\"]\n    .value_counts()\n)\n\nprint(\"\\nReport length:\")\ndisplay(\n    report_script_audit[\"report_words\"].describe()\n)\n\nreport_script_audit.to_csv(\n    OUTPUT_DIR / \"report_script_audit.csv\",\n    index=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:33.285715Z","iopub.execute_input":"2026-08-11T07:54:33.286066Z","iopub.status.idle":"2026-08-11T07:54:34.64945Z","shell.execute_reply.started":"2026-08-11T07:54:33.28604Z","shell.execute_reply":"2026-08-11T07:54:34.64848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 15: TEST DATASET STRUCTURE AUDIT\n# ============================================================\n\nprint(\"=\" * 70)\nprint(\"TEST DATASET AUDIT\")\nprint(\"=\" * 70)\n\nprint(\"\\nTEST CSV\")\nprint(\"Rows:\", len(test))\nprint(\"Unique studies:\", test[\"StudyInstanceUID\"].nunique())\n\nprint(\"\\nTEST SERIES CSV\")\nprint(\"Rows:\", len(test_series))\nprint(\n    \"Unique studies:\",\n    test_series[\"StudyInstanceUID\"].nunique()\n)\nprint(\n    \"Unique series:\",\n    test_series[\"SeriesInstanceUID\"].nunique()\n)\n\nprint(\"\\nSAMPLE SUBMISSION\")\nprint(\"Rows:\", len(sample_submission))\nprint(\n    \"Unique StudyInstanceUID:\",\n    sample_submission[\"StudyInstanceUID\"].nunique()\n)\n\nprint(\"\\nTest series per study:\")\n\ntest_series_per_study = (\n    test_series\n    .groupby(\"StudyInstanceUID\")\n    [\"SeriesInstanceUID\"]\n    .nunique()\n)\n\ndisplay(test_series_per_study)\n\nprint(\"\\nTest series metadata:\")\ndisplay(test_series)\n\nprint(\"\\nSample submission columns:\")\nprint(sample_submission.columns.tolist())\n\n# ------------------------------------------------------------\n# Check that sample submission IDs exactly match test IDs\n# ------------------------------------------------------------\n\ntest_ids = set(\n    test[\"StudyInstanceUID\"].astype(str)\n)\n\nsubmission_ids = set(\n    sample_submission[\"StudyInstanceUID\"].astype(str)\n)\n\nprint(\"\\nTest/submission ID comparison:\")\nprint(\"IDs in test but not submission:\",\n      test_ids - submission_ids)\n\nprint(\"IDs in submission but not test:\",\n      submission_ids - test_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:34.650726Z","iopub.execute_input":"2026-08-11T07:54:34.65119Z","iopub.status.idle":"2026-08-11T07:54:34.678128Z","shell.execute_reply.started":"2026-08-11T07:54:34.651141Z","shell.execute_reply":"2026-08-11T07:54:34.677055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 16: DICOM DIRECTORY STRUCTURE AUDIT\n# ============================================================\n\nfrom pathlib import Path\nfrom collections import Counter\nimport os\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"DICOM DIRECTORY STRUCTURE AUDIT\")\nprint(\"=\" * 70)\n\n# Use paths discovered earlier\nprint(\"\\nTrain series directory:\")\nprint(TRAIN_SERIES_DIR)\n\nprint(\"\\nTest series directory:\")\nprint(TEST_SERIES_DIR)\n\n# Basic existence checks\nprint(\"\\nDirectory existence:\")\nprint(\"Train series exists:\", TRAIN_SERIES_DIR.exists())\nprint(\"Test series exists:\", TEST_SERIES_DIR.exists())\n\n# ------------------------------------------------------------\n# Study directories\n# ------------------------------------------------------------\n\ntrain_studies = [\n    p for p in TRAIN_SERIES_DIR.iterdir()\n    if p.is_dir()\n]\n\ntest_studies = [\n    p for p in TEST_SERIES_DIR.iterdir()\n    if p.is_dir()\n]\n\nprint(\"\\nTrain study directories:\", len(train_studies))\nprint(\"Test study directories:\", len(test_studies))\n\n# ------------------------------------------------------------\n# Series directory counts\n# ------------------------------------------------------------\n\ntrain_series_dirs = []\ntest_series_dirs = []\n\nfor study_dir in train_studies:\n    train_series_dirs.extend(\n        [p for p in study_dir.iterdir() if p.is_dir()]\n    )\n\nfor study_dir in test_studies:\n    test_series_dirs.extend(\n        [p for p in study_dir.iterdir() if p.is_dir()]\n    )\n\nprint(\"\\nTrain series directories:\", len(train_series_dirs))\nprint(\"Test series directories:\", len(test_series_dirs))\n\n# ------------------------------------------------------------\n# Compare directory counts with CSV metadata\n# ------------------------------------------------------------\n\nprint(\"\\nCSV vs filesystem comparison:\")\n\nprint(\n    \"train_series.csv unique series:\",\n    train_series[\"SeriesInstanceUID\"].nunique()\n)\n\nprint(\n    \"Filesystem train series:\",\n    len(train_series_dirs)\n)\n\nprint(\n    \"test_series.csv unique series:\",\n    test_series[\"SeriesInstanceUID\"].nunique()\n)\n\nprint(\n    \"Filesystem test series:\",\n    len(test_series_dirs)\n)\n\n# ------------------------------------------------------------\n# Sample directory structures\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SAMPLE TRAIN DIRECTORY STRUCTURE\")\nprint(\"=\" * 70)\n\nfor study_dir in train_studies[:3]:\n\n    print(\"\\nStudy:\", study_dir.name)\n\n    series_dirs = [\n        p for p in study_dir.iterdir()\n        if p.is_dir()\n    ]\n\n    for series_dir in series_dirs[:5]:\n\n        print(\n            \"  Series:\",\n            series_dir.name\n        )\n\n        files = [\n            p for p in series_dir.iterdir()\n            if p.is_file()\n        ]\n\n        print(\n            \"    Files:\",\n            len(files)\n        )\n\n        if files:\n            print(\n                \"    Example file:\",\n                files[0].name\n            )\n\n# ------------------------------------------------------------\n# Sample test directory structures\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SAMPLE TEST DIRECTORY STRUCTURE\")\nprint(\"=\" * 70)\n\nfor study_dir in test_studies[:3]:\n\n    print(\"\\nStudy:\", study_dir.name)\n\n    series_dirs = [\n        p for p in study_dir.iterdir()\n        if p.is_dir()\n    ]\n\n    for series_dir in series_dirs:\n\n        files = [\n            p for p in series_dir.iterdir()\n            if p.is_file()\n        ]\n\n        print(\n            \"  Series:\",\n            series_dir.name,\n            \"| Files:\",\n            len(files)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:34.679526Z","iopub.execute_input":"2026-08-11T07:54:34.679847Z","iopub.status.idle":"2026-08-11T07:54:48.044552Z","shell.execute_reply.started":"2026-08-11T07:54:34.679811Z","shell.execute_reply":"2026-08-11T07:54:48.043428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 17: SLICE COUNT AUDIT\n# ============================================================\n\nfrom collections import Counter\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"SLICE COUNT AUDIT\")\nprint(\"=\" * 70)\n\n\ndef count_files_in_series(series_dir):\n    \"\"\"\n    Count files in one series directory without reading\n    DICOM pixel data.\n    \"\"\"\n    return sum(\n        1\n        for p in series_dir.iterdir()\n        if p.is_file()\n    )\n\n\n# ------------------------------------------------------------\n# TRAIN SERIES\n# ------------------------------------------------------------\n\nprint(\"\\nCounting train-series files...\")\n\ntrain_slice_rows = []\n\nfor i, series_dir in enumerate(train_series_dirs):\n\n    count = count_files_in_series(series_dir)\n\n    train_slice_rows.append({\n        \"StudyInstanceUID\": series_dir.parent.name,\n        \"SeriesInstanceUID\": series_dir.name,\n        \"slice_count\": count\n    })\n\n    if (i + 1) % 5000 == 0:\n        print(\n            f\"Processed {i + 1:,} / \"\n            f\"{len(train_series_dirs):,} train series\"\n        )\n\n\ntrain_slice_counts = pd.DataFrame(train_slice_rows)\n\n\n# ------------------------------------------------------------\n# TEST SERIES\n# ------------------------------------------------------------\n\nprint(\"\\nCounting test-series files...\")\n\ntest_slice_rows = []\n\nfor series_dir in test_series_dirs:\n\n    count = count_files_in_series(series_dir)\n\n    test_slice_rows.append({\n        \"StudyInstanceUID\": series_dir.parent.name,\n        \"SeriesInstanceUID\": series_dir.name,\n        \"slice_count\": count\n    })\n\n\ntest_slice_counts = pd.DataFrame(test_slice_rows)\n\n\n# ------------------------------------------------------------\n# BASIC TRAIN STATISTICS\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAIN SLICE COUNT STATISTICS\")\nprint(\"=\" * 70)\n\ndisplay(\n    train_slice_counts[\"slice_count\"].describe(\n        percentiles=[\n            0.01,\n            0.05,\n            0.10,\n            0.25,\n            0.50,\n            0.75,\n            0.90,\n            0.95,\n            0.99\n        ]\n    )\n)\n\n\n# ------------------------------------------------------------\n# TEST STATISTICS\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TEST SLICE COUNT STATISTICS\")\nprint(\"=\" * 70)\n\ndisplay(\n    test_slice_counts[\"slice_count\"].describe()\n)\n\n\n# ------------------------------------------------------------\n# SLICE COUNT FREQUENCY\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MOST COMMON TRAIN SLICE COUNTS\")\nprint(\"=\" * 70)\n\ndisplay(\n    train_slice_counts[\"slice_count\"]\n    .value_counts()\n    .sort_index()\n    .head(50)\n)\n\n\n# ------------------------------------------------------------\n# EXTREME SERIES\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SMALLEST TRAIN SERIES\")\nprint(\"=\" * 70)\n\ndisplay(\n    train_slice_counts\n    .sort_values(\"slice_count\")\n    .head(20)\n)\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LARGEST TRAIN SERIES\")\nprint(\"=\" * 70)\n\ndisplay(\n    train_slice_counts\n    .sort_values(\"slice_count\", ascending=False)\n    .head(20)\n)\n\n\n# ------------------------------------------------------------\n# Merge with series metadata\n# ------------------------------------------------------------\n\ntrain_slice_audit = train_slice_counts.merge(\n    train_series,\n    on=[\"StudyInstanceUID\", \"SeriesInstanceUID\"],\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\ntest_slice_audit = test_slice_counts.merge(\n    test_series,\n    on=[\"StudyInstanceUID\", \"SeriesInstanceUID\"],\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\n\n# ------------------------------------------------------------\n# Slice counts by anatomical plane\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SLICE COUNT BY ANATOMICAL PLANE\")\nprint(\"=\" * 70)\n\nplane_stats = (\n    train_slice_audit\n    .groupby(\"Anatomical_Plane\")[\"slice_count\"]\n    .agg(\n        [\"count\", \"min\", \"median\", \"mean\", \"max\"]\n    )\n    .sort_values(\"count\", ascending=False)\n)\n\ndisplay(plane_stats)\n\n\n# ------------------------------------------------------------\n# Slice counts by sequence properties\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SLICE COUNT BY SEQUENCE PROPERTIES\")\nprint(\"=\" * 70)\n\nsequence_stats = (\n    train_slice_audit\n    .groupby(\n        [\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\"\n        ]\n    )[\"slice_count\"]\n    .agg(\n        [\"count\", \"min\", \"median\", \"mean\", \"max\"]\n    )\n    .sort_values(\"count\", ascending=False)\n)\n\ndisplay(sequence_stats)\n\n\n# ------------------------------------------------------------\n# Save audit files\n# ------------------------------------------------------------\n\ntrain_slice_audit.to_csv(\n    OUTPUT_DIR / \"train_slice_audit.csv\",\n    index=False\n)\n\ntest_slice_audit.to_csv(\n    OUTPUT_DIR / \"test_slice_audit.csv\",\n    index=False\n)\n\nplane_stats.to_csv(\n    OUTPUT_DIR / \"slice_count_by_plane.csv\"\n)\n\nsequence_stats.to_csv(\n    OUTPUT_DIR / \"slice_count_by_sequence.csv\"\n)\n\n\nprint(\"\\nSaved:\")\nprint(\n    OUTPUT_DIR / \"train_slice_audit.csv\"\n)\nprint(\n    OUTPUT_DIR / \"test_slice_audit.csv\"\n)\nprint(\n    OUTPUT_DIR / \"slice_count_by_plane.csv\"\n)\nprint(\n    OUTPUT_DIR / \"slice_count_by_sequence.csv\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T07:54:48.046057Z","iopub.execute_input":"2026-08-11T07:54:48.046431Z","iopub.status.idle":"2026-08-11T08:12:20.196519Z","shell.execute_reply.started":"2026-08-11T07:54:48.046392Z","shell.execute_reply":"2026-08-11T08:12:20.19514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 18: REPRESENTATIVE DICOM METADATA AUDIT\n# ============================================================\n\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom pathlib import Path\n\n\nprint(\"=\" * 70)\nprint(\"REPRESENTATIVE DICOM METADATA AUDIT\")\nprint(\"=\" * 70)\n\n\n# ------------------------------------------------------------\n# 1. Verify pydicom\n# ------------------------------------------------------------\n\nprint(\"\\nPydicom version:\")\nprint(pydicom.__version__)\n\n\n# ------------------------------------------------------------\n# 2. Helper function\n# ------------------------------------------------------------\n\ndef safe_get(ds, keyword, default=None):\n    \"\"\"\n    Safely retrieve a DICOM attribute.\n    \"\"\"\n    try:\n        value = getattr(ds, keyword, default)\n\n        if value is None:\n            return default\n\n        return value\n\n    except Exception:\n        return default\n\n\ndef normalize_dicom_value(value):\n    \"\"\"\n    Convert DICOM values into CSV-friendly representations.\n    \"\"\"\n\n    if value is None:\n        return None\n\n    if isinstance(value, (list, tuple)):\n        return \"|\".join(str(x) for x in value)\n\n    return str(value)\n\n\n# ------------------------------------------------------------\n# 3. Select representative series\n# ------------------------------------------------------------\n\nmetadata_source = train_slice_audit.copy()\n\n# Keep only rows with usable series metadata.\nmetadata_source = metadata_source[\n    metadata_source[\"SeriesInstanceUID\"].notna()\n].copy()\n\n\nselected_series = []\n\n\n# ------------------------------------------------------------\n# A. One representative series from each plane\n# ------------------------------------------------------------\n\nfor plane in sorted(\n    metadata_source[\"Anatomical_Plane\"]\n    .dropna()\n    .astype(str)\n    .unique()\n):\n\n    candidates = metadata_source[\n        metadata_source[\"Anatomical_Plane\"].astype(str) == plane\n    ]\n\n    if len(candidates) > 0:\n\n        selected_series.append(\n            candidates.sample(\n                n=1,\n                random_state=42\n            ).iloc[0]\n        )\n\n\n# ------------------------------------------------------------\n# B. One representative series from each\n#    sequence-property combination\n# ------------------------------------------------------------\n\nsequence_columns = [\n    \"Anatomical_Plane\",\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\"\n]\n\nfor _, group in metadata_source.groupby(\n    sequence_columns,\n    dropna=False\n):\n\n    if len(group) == 0:\n        continue\n\n    selected_series.append(\n        group.sample(\n            n=1,\n            random_state=42\n        ).iloc[0]\n    )\n\n\n# ------------------------------------------------------------\n# C. Very small series\n# ------------------------------------------------------------\n\nsmallest = (\n    metadata_source\n    .sort_values(\"slice_count\")\n    .head(2)\n)\n\nfor _, row in smallest.iterrows():\n    selected_series.append(row)\n\n\n# ------------------------------------------------------------\n# D. Very large series\n# ------------------------------------------------------------\n\nlargest = (\n    metadata_source\n    .sort_values(\n        \"slice_count\",\n        ascending=False\n    )\n    .head(3)\n)\n\nfor _, row in largest.iterrows():\n    selected_series.append(row)\n\n\n# ------------------------------------------------------------\n# E. A few labeled studies if available\n# ------------------------------------------------------------\n\nlabeled_ids = set(\n    labeled_train[\"StudyInstanceUID\"].astype(str)\n)\n\nlabeled_series = metadata_source[\n    metadata_source[\"StudyInstanceUID\"]\n    .astype(str)\n    .isin(labeled_ids)\n]\n\nif len(labeled_series) > 0:\n\n    selected_series.append(\n        labeled_series.sample(\n            n=1,\n            random_state=42\n        ).iloc[0]\n    )\n\n\n# ------------------------------------------------------------\n# Remove duplicate series\n# ------------------------------------------------------------\n\nselected_df = pd.DataFrame(selected_series)\n\nselected_df = (\n    selected_df\n    .drop_duplicates(\n        subset=[\"SeriesInstanceUID\"]\n    )\n    .reset_index(drop=True)\n)\n\n\nprint(\"\\nSelected representative series:\")\nprint(len(selected_df))\n\ndisplay(\n    selected_df[\n        [\n            \"StudyInstanceUID\",\n            \"SeriesInstanceUID\",\n            \"slice_count\",\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\"\n        ]\n    ]\n)\n\n\n# ------------------------------------------------------------\n# 4. Read representative DICOM files\n# ------------------------------------------------------------\n\nmetadata_rows = []\nfailed_files = []\n\n\nfor _, series_row in selected_df.iterrows():\n\n    study_uid = str(\n        series_row[\"StudyInstanceUID\"]\n    )\n\n    series_uid = str(\n        series_row[\"SeriesInstanceUID\"]\n    )\n\n    series_dir = (\n        TRAIN_SERIES_DIR\n        / study_uid\n        / series_uid\n    )\n\n    files = sorted(\n        [\n            p\n            for p in series_dir.iterdir()\n            if p.is_file()\n        ]\n    )\n\n    if not files:\n        print(\n            \"\\nWARNING: No files found:\",\n            series_dir\n        )\n        continue\n\n\n    # Select first, middle and last files.\n    positions = sorted(\n        set(\n            [\n                0,\n                len(files) // 2,\n                len(files) - 1\n            ]\n        )\n    )\n\n\n    for position in positions:\n\n        filepath = files[position]\n\n        try:\n\n            # IMPORTANT:\n            # stop_before_pixels=True means pixel arrays\n            # are NOT loaded.\n            ds = pydicom.dcmread(\n                filepath,\n                stop_before_pixels=True,\n                force=True\n            )\n\n\n            row = {\n                \"StudyInstanceUID\":\n                    study_uid,\n\n                \"SeriesInstanceUID\":\n                    series_uid,\n\n                \"FileName\":\n                    filepath.name,\n\n                \"FilePosition\":\n                    position,\n\n                \"DirectorySliceCount\":\n                    len(files),\n\n                \"Rows\":\n                    safe_get(ds, \"Rows\"),\n\n                \"Columns\":\n                    safe_get(ds, \"Columns\"),\n\n                \"PixelSpacing\":\n                    normalize_dicom_value(\n                        safe_get(ds, \"PixelSpacing\")\n                    ),\n\n                \"SliceThickness\":\n                    safe_get(ds, \"SliceThickness\"),\n\n                \"SpacingBetweenSlices\":\n                    safe_get(\n                        ds,\n                        \"SpacingBetweenSlices\"\n                    ),\n\n                \"ImageOrientationPatient\":\n                    normalize_dicom_value(\n                        safe_get(\n                            ds,\n                            \"ImageOrientationPatient\"\n                        )\n                    ),\n\n                \"ImagePositionPatient\":\n                    normalize_dicom_value(\n                        safe_get(\n                            ds,\n                            \"ImagePositionPatient\"\n                        )\n                    ),\n\n                \"InstanceNumber\":\n                    safe_get(ds, \"InstanceNumber\"),\n\n                \"PhotometricInterpretation\":\n                    safe_get(\n                        ds,\n                        \"PhotometricInterpretation\"\n                    ),\n\n                \"BitsAllocated\":\n                    safe_get(\n                        ds,\n                        \"BitsAllocated\"\n                    ),\n\n                \"BitsStored\":\n                    safe_get(\n                        ds,\n                        \"BitsStored\"\n                    ),\n\n                \"HighBit\":\n                    safe_get(ds, \"HighBit\"),\n\n                \"PixelRepresentation\":\n                    safe_get(\n                        ds,\n                        \"PixelRepresentation\"\n                    ),\n\n                \"RescaleSlope\":\n                    safe_get(\n                        ds,\n                        \"RescaleSlope\"\n                    ),\n\n                \"RescaleIntercept\":\n                    safe_get(\n                        ds,\n                        \"RescaleIntercept\"\n                    ),\n\n                \"TransferSyntaxUID\":\n                    normalize_dicom_value(\n                        safe_get(\n                            ds,\n                            \"file_meta\"\n                        )\n                    ),\n\n                \"SOPInstanceUID\":\n                    safe_get(\n                        ds,\n                        \"SOPInstanceUID\"\n                    ),\n\n                \"DICOMStudyInstanceUID\":\n                    safe_get(\n                        ds,\n                        \"StudyInstanceUID\"\n                    ),\n\n                \"DICOMSeriesInstanceUID\":\n                    safe_get(\n                        ds,\n                        \"SeriesInstanceUID\"\n                    ),\n\n                \"DICOMModality\":\n                    safe_get(ds, \"Modality\"),\n\n                \"Manufacturer\":\n                    safe_get(ds, \"Manufacturer\"),\n\n                \"ManufacturerModelName\":\n                    safe_get(\n                        ds,\n                        \"ManufacturerModelName\"\n                    ),\n\n                \"SeriesDescription\":\n                    safe_get(\n                        ds,\n                        \"SeriesDescription\"\n                    )\n            }\n\n            metadata_rows.append(row)\n\n\n        except Exception as e:\n\n            failed_files.append({\n                \"StudyInstanceUID\":\n                    study_uid,\n\n                \"SeriesInstanceUID\":\n                    series_uid,\n\n                \"FileName\":\n                    filepath.name,\n\n                \"Error\":\n                    repr(e)\n            })\n\n\n# ------------------------------------------------------------\n# 5. Create metadata DataFrame\n# ------------------------------------------------------------\n\ndicom_metadata_audit = pd.DataFrame(\n    metadata_rows\n)\n\ndicom_failures = pd.DataFrame(\n    failed_files\n)\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"DICOM METADATA RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    \"\\nDICOM files successfully inspected:\",\n    len(dicom_metadata_audit)\n)\n\nprint(\n    \"DICOM files that failed:\",\n    len(dicom_failures)\n)\n\n\n# ------------------------------------------------------------\n# 6. Display metadata\n# ------------------------------------------------------------\n\ndisplay(\n    dicom_metadata_audit\n)\n\n\n# ------------------------------------------------------------\n# 7. Missing metadata analysis\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MISSING METADATA\")\nprint(\"=\" * 70)\n\nif len(dicom_metadata_audit) > 0:\n\n    metadata_missingness = (\n        dicom_metadata_audit\n        .isna()\n        .sum()\n        .sort_values(\n            ascending=False\n        )\n    )\n\n    display(\n        metadata_missingness[\n            metadata_missingness > 0\n        ]\n    )\n\n\n# ------------------------------------------------------------\n# 8. Common dimensions\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"IMAGE DIMENSIONS\")\nprint(\"=\" * 70)\n\nif len(dicom_metadata_audit) > 0:\n\n    dimension_counts = (\n        dicom_metadata_audit[\n            [\"Rows\", \"Columns\"]\n        ]\n        .value_counts()\n        .reset_index(\n            name=\"count\"\n        )\n    )\n\n    display(\n        dimension_counts.head(20)\n    )\n\n\n# ------------------------------------------------------------\n# 9. Photometric interpretation\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PHOTOMETRIC INTERPRETATION\")\nprint(\"=\" * 70)\n\ndisplay(\n    dicom_metadata_audit[\n        \"PhotometricInterpretation\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\n# ------------------------------------------------------------\n# 10. Transfer syntax\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRANSFER SYNTAX\")\nprint(\"=\" * 70)\n\ndisplay(\n    dicom_metadata_audit[\n        \"TransferSyntaxUID\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\n# ------------------------------------------------------------\n# 11. Modality\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MODALITY\")\nprint(\"=\" * 70)\n\ndisplay(\n    dicom_metadata_audit[\n        \"DICOMModality\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\n# ------------------------------------------------------------\n# 12. Save results\n# ------------------------------------------------------------\n\ndicom_metadata_audit.to_csv(\n    OUTPUT_DIR / \"dicom_metadata_audit.csv\",\n    index=False\n)\n\ndicom_failures.to_csv(\n    OUTPUT_DIR / \"dicom_metadata_failures.csv\",\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(\n    OUTPUT_DIR / \"dicom_metadata_audit.csv\"\n)\nprint(\n    OUTPUT_DIR / \"dicom_metadata_failures.csv\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:12:20.19816Z","iopub.execute_input":"2026-08-11T08:12:20.198456Z","iopub.status.idle":"2026-08-11T08:12:21.965854Z","shell.execute_reply.started":"2026-08-11T08:12:20.198432Z","shell.execute_reply":"2026-08-11T08:12:21.964628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 19: VISUAL DICOM + SLICE ORDER AUDIT\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\n\nfrom pathlib import Path\n\nprint(\"=\" * 70)\nprint(\"VISUAL DICOM + SLICE ORDER AUDIT\")\nprint(\"=\" * 70)\n\nTRAIN_SERIES_DIR = Path(\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series\"\n)\n\nAUDIT_DIR = Path(\"/kaggle/working/rsna_knee_audit\")\nAUDIT_DIR.mkdir(parents=True, exist_ok=True)\n\n\n# ------------------------------------------------------------\n# 1. Helper: read a series\n# ------------------------------------------------------------\n\ndef read_dicom_series(series_path):\n    records = []\n\n    for file_path in sorted(series_path.glob(\"*.dcm\")):\n        try:\n            ds = pydicom.dcmread(str(file_path), force=False)\n\n            if not hasattr(ds, \"PixelData\"):\n                continue\n\n            position = getattr(ds, \"ImagePositionPatient\", None)\n            orientation = getattr(ds, \"ImageOrientationPatient\", None)\n\n            if position is not None:\n                position = np.asarray(position, dtype=np.float64)\n            else:\n                position = None\n\n            if orientation is not None:\n                orientation = np.asarray(orientation, dtype=np.float64)\n            else:\n                orientation = None\n\n            records.append({\n                \"path\": file_path,\n                \"ds\": ds,\n                \"position\": position,\n                \"orientation\": orientation,\n                \"instance_number\": getattr(ds, \"InstanceNumber\", None)\n            })\n\n        except Exception:\n            continue\n\n    if not records:\n        raise ValueError(f\"No readable DICOM images found: {series_path}\")\n\n    # --------------------------------------------------------\n    # Determine spatial sorting direction\n    # --------------------------------------------------------\n\n    valid_positions = [\n        r[\"position\"] for r in records\n        if r[\"position\"] is not None\n    ]\n\n    orientation = None\n\n    for r in records:\n        if r[\"orientation\"] is not None:\n            orientation = r[\"orientation\"]\n            break\n\n    if len(valid_positions) == len(records) and orientation is not None:\n        row_cosines = orientation[:3]\n        col_cosines = orientation[3:]\n        normal = np.cross(row_cosines, col_cosines)\n\n        for r in records:\n            r[\"sort_position\"] = float(\n                np.dot(r[\"position\"], normal)\n            )\n\n        records.sort(key=lambda x: x[\"sort_position\"])\n\n        sort_method = \"ImagePositionPatient + orientation normal\"\n\n    elif all(r[\"position\"] is not None for r in records):\n        for r in records:\n            r[\"sort_position\"] = float(\n                r[\"position\"][2]\n            )\n\n        records.sort(key=lambda x: x[\"sort_position\"])\n\n        sort_method = \"ImagePositionPatient Z\"\n\n    elif all(r[\"instance_number\"] is not None for r in records):\n        records.sort(key=lambda x: x[\"instance_number\"])\n\n        sort_method = \"InstanceNumber\"\n\n    else:\n        records.sort(key=lambda x: str(x[\"path\"]))\n\n        sort_method = \"Filename fallback\"\n\n    images = []\n\n    for r in records:\n        try:\n            pixel_array = r[\"ds\"].pixel_array.astype(np.float32)\n            images.append(pixel_array)\n        except Exception:\n            continue\n\n    return records, images, sort_method\n\n\n# ------------------------------------------------------------\n# 2. Find representative series\n# ------------------------------------------------------------\n\nprint(\"\\nSearching for representative series...\")\n\nseries_dirs = [\n    p for p in TRAIN_SERIES_DIR.rglob(\"*\")\n    if p.is_dir() and any(p.glob(\"*.dcm\"))\n]\n\nprint(f\"Series directories discovered: {len(series_dirs)}\")\n\n\n# ------------------------------------------------------------\n# 3. Gather basic information without reading every pixel\n# ------------------------------------------------------------\n\nrepresentatives = {\n    \"sagittal_normal\": None,\n    \"coronal_normal\": None,\n    \"axial_normal\": None,\n    \"small_11\": None,\n    \"large_320\": None\n}\n\nfor series_path in series_dirs:\n\n    files = list(series_path.glob(\"*.dcm\"))\n    count = len(files)\n\n    if count == 0:\n        continue\n\n    try:\n        ds = pydicom.dcmread(\n            str(files[0]),\n            stop_before_pixels=True\n        )\n    except Exception:\n        continue\n\n    plane = str(getattr(ds, \"AnatomicalPlane\", \"\")).strip()\n\n    if plane == \"\":\n        description = str(\n            getattr(ds, \"SeriesDescription\", \"\")\n        ).upper()\n\n        if \"SAG\" in description:\n            plane = \"Sagittal\"\n        elif \"COR\" in description:\n            plane = \"Coronal\"\n        elif \"TRA\" in description or \"AX\" in description:\n            plane = \"Axial\"\n\n    if count == 11 and representatives[\"small_11\"] is None:\n        representatives[\"small_11\"] = series_path\n\n    if count == 320 and representatives[\"large_320\"] is None:\n        representatives[\"large_320\"] = series_path\n\n    if 20 <= count <= 40:\n\n        if (\n            plane == \"Sagittal\"\n            and representatives[\"sagittal_normal\"] is None\n        ):\n            representatives[\"sagittal_normal\"] = series_path\n\n        elif (\n            plane == \"Coronal\"\n            and representatives[\"coronal_normal\"] is None\n        ):\n            representatives[\"coronal_normal\"] = series_path\n\n        elif (\n            plane == \"Axial\"\n            and representatives[\"axial_normal\"] is None\n        ):\n            representatives[\"axial_normal\"] = series_path\n\n    if all(v is not None for v in representatives.values()):\n        break\n\n\nprint(\"\\nSelected representatives:\")\n\nfor name, path in representatives.items():\n    print(f\"{name}: {path}\")\n\n\n# ------------------------------------------------------------\n# 4. Analyze and visualize each representative\n# ------------------------------------------------------------\n\naudit_rows = []\n\nfor name, series_path in representatives.items():\n\n    if series_path is None:\n        print(f\"\\n{name}: NOT FOUND\")\n        continue\n\n    print(\"\\n\" + \"=\" * 70)\n    print(f\"SERIES: {name}\")\n    print(\"=\" * 70)\n\n    records, images, sort_method = read_dicom_series(series_path)\n\n    if not images:\n        print(\"No pixel arrays could be decoded.\")\n        continue\n\n    first_ds = records[0][\"ds\"]\n\n    volume = np.stack(images)\n\n    print(f\"Directory: {series_path}\")\n    print(f\"Slice count: {len(images)}\")\n    print(f\"Volume shape: {volume.shape}\")\n    print(f\"Dtype: {volume.dtype}\")\n    print(f\"Min: {volume.min():.3f}\")\n    print(f\"Max: {volume.max():.3f}\")\n    print(f\"Mean: {volume.mean():.3f}\")\n    print(f\"Std: {volume.std():.3f}\")\n    print(f\"Sort method: {sort_method}\")\n\n    print(\n        \"Pixel spacing:\",\n        getattr(first_ds, \"PixelSpacing\", \"MISSING\")\n    )\n\n    print(\n        \"Slice thickness:\",\n        getattr(first_ds, \"SliceThickness\", \"MISSING\")\n    )\n\n    print(\n        \"Spacing between slices:\",\n        getattr(first_ds, \"SpacingBetweenSlices\", \"MISSING\")\n    )\n\n    print(\n        \"Image orientation:\",\n        getattr(first_ds, \"ImageOrientationPatient\", \"MISSING\")\n    )\n\n    print(\n        \"Image position available:\",\n        all(\n            r[\"position\"] is not None\n            for r in records\n        )\n    )\n\n    print(\n        \"Photometric interpretation:\",\n        getattr(\n            first_ds,\n            \"PhotometricInterpretation\",\n            \"MISSING\"\n        )\n    )\n\n    print(\n        \"Series description:\",\n        getattr(\n            first_ds,\n            \"SeriesDescription\",\n            \"MISSING\"\n        )\n    )\n\n    # --------------------------------------------------------\n    # Check spatial positions\n    # --------------------------------------------------------\n\n    positions = [\n        r[\"sort_position\"]\n        for r in records\n        if \"sort_position\" in r\n    ]\n\n    duplicate_positions = 0\n\n    if positions:\n        rounded_positions = np.round(\n            np.asarray(positions),\n            decimals=5\n        )\n\n        duplicate_positions = (\n            len(rounded_positions)\n            - len(np.unique(rounded_positions))\n        )\n\n    print(\n        \"Duplicate spatial positions:\",\n        duplicate_positions\n    )\n\n    # --------------------------------------------------------\n    # Save audit row\n    # --------------------------------------------------------\n\n    audit_rows.append({\n        \"representative\": name,\n        \"series_path\": str(series_path),\n        \"slice_count\": len(images),\n        \"rows\": volume.shape[1],\n        \"columns\": volume.shape[2],\n        \"dtype\": str(volume.dtype),\n        \"min\": float(volume.min()),\n        \"max\": float(volume.max()),\n        \"mean\": float(volume.mean()),\n        \"std\": float(volume.std()),\n        \"sort_method\": sort_method,\n        \"pixel_spacing\": str(\n            getattr(first_ds, \"PixelSpacing\", \"MISSING\")\n        ),\n        \"slice_thickness\": str(\n            getattr(first_ds, \"SliceThickness\", \"MISSING\")\n        ),\n        \"spacing_between_slices\": str(\n            getattr(\n                first_ds,\n                \"SpacingBetweenSlices\",\n                \"MISSING\"\n            )\n        ),\n        \"photometric\": str(\n            getattr(\n                first_ds,\n                \"PhotometricInterpretation\",\n                \"MISSING\"\n            )\n        ),\n        \"duplicate_spatial_positions\": duplicate_positions\n    })\n\n    # --------------------------------------------------------\n    # Normalize only for visualization\n    # --------------------------------------------------------\n\n    display_volume = volume.copy()\n\n    low = np.percentile(display_volume, 1)\n    high = np.percentile(display_volume, 99)\n\n    if high > low:\n        display_volume = np.clip(\n            display_volume,\n            low,\n            high\n        )\n\n        display_volume = (\n            display_volume - low\n        ) / (high - low)\n\n    # --------------------------------------------------------\n    # First / middle / last\n    # --------------------------------------------------------\n\n    indices = [\n        0,\n        len(images) // 2,\n        len(images) - 1\n    ]\n\n    fig, axes = plt.subplots(\n        1,\n        3,\n        figsize=(15, 5)\n    )\n\n    for ax, idx in zip(axes, indices):\n\n        ax.imshow(\n            display_volume[idx],\n            cmap=\"gray\"\n        )\n\n        ax.set_title(\n            f\"Slice {idx + 1}/{len(images)}\"\n        )\n\n        ax.axis(\"off\")\n\n    fig.suptitle(\n        f\"{name}\\n{series_path.name}\",\n        fontsize=12\n    )\n\n    plt.tight_layout()\n\n    output_path = (\n        AUDIT_DIR /\n        f\"cell19_{name}_first_middle_last.png\"\n    )\n\n    plt.savefig(\n        output_path,\n        dpi=150,\n        bbox_inches=\"tight\"\n    )\n\n    plt.show()\n\n    # --------------------------------------------------------\n    # Contact sheet\n    # --------------------------------------------------------\n\n    contact_count = min(len(images), 12)\n\n    contact_indices = np.linspace(\n        0,\n        len(images) - 1,\n        contact_count,\n        dtype=int\n    )\n\n    fig, axes = plt.subplots(\n        3,\n        4,\n        figsize=(12, 9)\n    )\n\n    axes = axes.ravel()\n\n    for ax in axes:\n        ax.axis(\"off\")\n\n    for ax, idx in zip(\n        axes,\n        contact_indices\n    ):\n\n        ax.imshow(\n            display_volume[idx],\n            cmap=\"gray\"\n        )\n\n        ax.set_title(\n            f\"{idx + 1}/{len(images)}\"\n        )\n\n        ax.axis(\"off\")\n\n    fig.suptitle(\n        f\"Contact Sheet: {name}\",\n        fontsize=14\n    )\n\n    plt.tight_layout()\n\n    output_path = (\n        AUDIT_DIR /\n        f\"cell19_{name}_contact_sheet.png\"\n    )\n\n    plt.savefig(\n        output_path,\n        dpi=150,\n        bbox_inches=\"tight\"\n    )\n\n    plt.show()\n\n\n# ------------------------------------------------------------\n# 5. Save summary\n# ------------------------------------------------------------\n\naudit_df = pd.DataFrame(audit_rows)\n\nsummary_path = (\n    AUDIT_DIR /\n    \"cell19_visual_dicom_audit.csv\"\n)\n\naudit_df.to_csv(\n    summary_path,\n    index=False\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 19 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"\\nSaved:\")\nprint(summary_path)\n\nprint(\"\\nGenerated visualization files:\")\nfor p in sorted(AUDIT_DIR.glob(\"cell19_*.png\")):\n    print(p)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:12:21.967643Z","iopub.execute_input":"2026-08-11T08:12:21.968113Z","iopub.status.idle":"2026-08-11T08:30:51.624845Z","shell.execute_reply.started":"2026-08-11T08:12:21.968083Z","shell.execute_reply":"2026-08-11T08:30:51.623968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 20: BROAD DICOM METADATA + PIXEL AUDIT\n# ============================================================\n\nimport os\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\nfrom pathlib import Path\n\nprint(\"=\" * 70)\nprint(\"BROAD DICOM METADATA + PIXEL AUDIT\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------------------\n# Configuration\n# ------------------------------------------------------------\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\nAUDIT_DIR = Path(\"/kaggle/working/rsna_knee_audit\")\n\nTRAIN_SERIES_DIR = Path(\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series\"\n)\n\nTEST_SERIES_DIR = Path(\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/test_series\"\n)\n\nTRAIN_SERIES_CSV = (\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train_series.csv\"\n)\n\nTEST_SERIES_CSV = (\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/test_series.csv\"\n)\n\nOUTPUT_FILE = AUDIT_DIR / \"dicom_audit.csv\"\n\nTRAIN_SAMPLE_SIZE = 250\n\n\n# ------------------------------------------------------------\n# 1. Load series metadata\n# ------------------------------------------------------------\n\ntrain_series = pd.read_csv(TRAIN_SERIES_CSV)\ntest_series = pd.read_csv(TEST_SERIES_CSV)\n\nprint(\"\\nTrain series rows:\", len(train_series))\nprint(\"Test series rows:\", len(test_series))\n\n\n# ------------------------------------------------------------\n# 2. Build representative training sample\n# ------------------------------------------------------------\n\nprint(\"\\nCreating stratified training sample...\")\n\nsample_parts = []\n\n# Slice-count audit is available from Cell 17.\n# Use it if available, otherwise calculate counts from filesystem.\n\nslice_audit_file = AUDIT_DIR / \"train_slice_audit.csv\"\n\nif slice_audit_file.exists():\n\n    slice_audit = pd.read_csv(slice_audit_file)\n\n    train_meta = train_series.merge(\n        slice_audit[\n            [\n                \"StudyInstanceUID\",\n                \"SeriesInstanceUID\",\n                \"slice_count\"\n            ]\n        ],\n        on=[\"StudyInstanceUID\", \"SeriesInstanceUID\"],\n        how=\"left\"\n    )\n\nelse:\n\n    print(\n        \"train_slice_audit.csv not found. \"\n        \"Using metadata-only sampling.\"\n    )\n\n    train_meta = train_series.copy()\n    train_meta[\"slice_count\"] = np.nan\n\n\n# ------------------------------------------------------------\n# Define sampling groups\n# ------------------------------------------------------------\n\ntrain_meta[\"plane_group\"] = (\n    train_meta[\"Anatomical_Plane\"]\n    .fillna(\"UNKNOWN\")\n    .astype(str)\n)\n\ntrain_meta[\"sequence_group\"] = (\n    train_meta[\"Fluid_Sensitive\"].fillna(-1).astype(str)\n    + \"_\"\n    + train_meta[\"Fat_Suppression\"].fillna(-1).astype(str)\n)\n\nif train_meta[\"slice_count\"].notna().any():\n\n    train_meta[\"slice_group\"] = pd.cut(\n        train_meta[\"slice_count\"],\n        bins=[0, 15, 25, 35, 50, 100, 1000],\n        labels=[\n            \"11-15\",\n            \"16-25\",\n            \"26-35\",\n            \"36-50\",\n            \"51-100\",\n            \"101+\"\n        ],\n        include_lowest=True\n    )\n\nelse:\n\n    train_meta[\"slice_group\"] = \"UNKNOWN\"\n\n\n# ------------------------------------------------------------\n# Stratified sample\n# ------------------------------------------------------------\n\ngrouped = train_meta.groupby(\n    [\n        \"plane_group\",\n        \"sequence_group\",\n        \"slice_group\"\n    ],\n    dropna=False,\n    observed=True\n)\n\nn_groups = len(grouped)\n\nper_group = max(\n    1,\n    TRAIN_SAMPLE_SIZE // max(n_groups, 1)\n)\n\nfor _, group in grouped:\n\n    take = min(\n        len(group),\n        per_group\n    )\n\n    if take > 0:\n        sample_parts.append(\n            group.sample(\n                n=take,\n                random_state=SEED\n            )\n        )\n\ntrain_sample = pd.concat(\n    sample_parts,\n    ignore_index=True\n) if sample_parts else train_meta.sample(\n    n=min(TRAIN_SAMPLE_SIZE, len(train_meta)),\n    random_state=SEED\n)\n\n# Fill remaining slots if stratified sampling produced too few rows.\n\nif len(train_sample) < TRAIN_SAMPLE_SIZE:\n\n    remaining = train_meta[\n        ~train_meta[\"SeriesInstanceUID\"].isin(\n            train_sample[\"SeriesInstanceUID\"]\n        )\n    ]\n\n    extra_n = min(\n        TRAIN_SAMPLE_SIZE - len(train_sample),\n        len(remaining)\n    )\n\n    if extra_n > 0:\n\n        extra = remaining.sample(\n            n=extra_n,\n            random_state=SEED\n        )\n\n        train_sample = pd.concat(\n            [\n                train_sample,\n                extra\n            ],\n            ignore_index=True\n        )\n\ntrain_sample = train_sample.drop_duplicates(\n    subset=[\"SeriesInstanceUID\"]\n).head(TRAIN_SAMPLE_SIZE)\n\nprint(\n    f\"Selected {len(train_sample)} \"\n    \"training series for detailed DICOM audit.\"\n)\n\n\n# ------------------------------------------------------------\n# 3. Helper functions\n# ------------------------------------------------------------\n\ndef get_series_directory(root, study_uid, series_uid):\n\n    return (\n        Path(root)\n        / str(study_uid)\n        / str(series_uid)\n    )\n\n\ndef choose_middle_file(series_dir):\n\n    files = sorted(\n        series_dir.glob(\"*.dcm\")\n    )\n\n    if not files:\n        return None\n\n    return files[len(files) // 2]\n\n\ndef safe_value(ds, field):\n\n    try:\n        value = getattr(ds, field, None)\n\n        if value is None:\n            return None\n\n        if isinstance(value, (list, tuple)):\n            return str(list(value))\n\n        return str(value)\n\n    except Exception:\n\n        return None\n\n\ndef safe_float(ds, field):\n\n    try:\n\n        value = getattr(ds, field, None)\n\n        if value is None:\n            return np.nan\n\n        return float(value)\n\n    except Exception:\n\n        return np.nan\n\n\ndef get_pixel_statistics(ds):\n\n    try:\n\n        if not hasattr(ds, \"PixelData\"):\n            return {\n                \"pixel_loaded\": False,\n                \"pixel_min\": np.nan,\n                \"pixel_max\": np.nan,\n                \"pixel_mean\": np.nan,\n                \"pixel_std\": np.nan,\n                \"pixel_p01\": np.nan,\n                \"pixel_p50\": np.nan,\n                \"pixel_p99\": np.nan\n            }\n\n        arr = ds.pixel_array.astype(\n            np.float32\n        )\n\n        # Apply rescale when present.\n\n        slope = safe_float(\n            ds,\n            \"RescaleSlope\"\n        )\n\n        intercept = safe_float(\n            ds,\n            \"RescaleIntercept\"\n        )\n\n        if np.isfinite(slope):\n            arr = arr * slope\n\n        if np.isfinite(intercept):\n            arr = arr + intercept\n\n        return {\n            \"pixel_loaded\": True,\n            \"pixel_min\": float(np.min(arr)),\n            \"pixel_max\": float(np.max(arr)),\n            \"pixel_mean\": float(np.mean(arr)),\n            \"pixel_std\": float(np.std(arr)),\n            \"pixel_p01\": float(np.percentile(arr, 1)),\n            \"pixel_p50\": float(np.percentile(arr, 50)),\n            \"pixel_p99\": float(np.percentile(arr, 99))\n        }\n\n    except Exception:\n\n        return {\n            \"pixel_loaded\": False,\n            \"pixel_min\": np.nan,\n            \"pixel_max\": np.nan,\n            \"pixel_mean\": np.nan,\n            \"pixel_std\": np.nan,\n            \"pixel_p01\": np.nan,\n            \"pixel_p50\": np.nan,\n            \"pixel_p99\": np.nan\n        }\n\n\n# ------------------------------------------------------------\n# 4. Audit function\n# ------------------------------------------------------------\n\naudit_rows = []\n\nloaded_count = 0\nfailed_count = 0\nmissing_pixel_count = 0\n\n\ndef audit_series(\n    root,\n    study_uid,\n    series_uid,\n    source,\n    metadata_row=None\n):\n\n    global loaded_count\n    global failed_count\n    global missing_pixel_count\n\n    series_dir = get_series_directory(\n        root,\n        study_uid,\n        series_uid\n    )\n\n    dcm_file = choose_middle_file(\n        series_dir\n    )\n\n    if dcm_file is None:\n\n        failed_count += 1\n\n        return {\n            \"source\": source,\n            \"StudyInstanceUID\": study_uid,\n            \"SeriesInstanceUID\": series_uid,\n            \"file_path\": str(series_dir),\n            \"read_status\": \"NO_DICOM_FILE\"\n        }\n\n    try:\n\n        ds = pydicom.dcmread(\n            str(dcm_file),\n            force=False\n        )\n\n        loaded_count += 1\n\n    except Exception as e:\n\n        failed_count += 1\n\n        return {\n            \"source\": source,\n            \"StudyInstanceUID\": study_uid,\n            \"SeriesInstanceUID\": series_uid,\n            \"file_path\": str(dcm_file),\n            \"read_status\": \"READ_FAILED\",\n            \"error\": str(e)[:300]\n        }\n\n\n    if not hasattr(ds, \"PixelData\"):\n        missing_pixel_count += 1\n\n\n    row = {\n        \"source\": source,\n\n        \"StudyInstanceUID\": str(\n            study_uid\n        ),\n\n        \"SeriesInstanceUID\": str(\n            series_uid\n        ),\n\n        \"file_path\": str(\n            dcm_file\n        ),\n\n        \"read_status\": \"OK\",\n\n        \"Modality\": safe_value(\n            ds,\n            \"Modality\"\n        ),\n\n        \"SOPClassUID\": safe_value(\n            ds,\n            \"SOPClassUID\"\n        ),\n\n        \"TransferSyntaxUID\": safe_value(\n            getattr(\n                ds,\n                \"file_meta\",\n                None\n            ),\n            \"TransferSyntaxUID\"\n        ) if hasattr(\n            ds,\n            \"file_meta\"\n        ) else None,\n\n        \"Rows\": safe_value(\n            ds,\n            \"Rows\"\n        ),\n\n        \"Columns\": safe_value(\n            ds,\n            \"Columns\"\n        ),\n\n        \"SamplesPerPixel\": safe_value(\n            ds,\n            \"SamplesPerPixel\"\n        ),\n\n        \"PhotometricInterpretation\": safe_value(\n            ds,\n            \"PhotometricInterpretation\"\n        ),\n\n        \"BitsAllocated\": safe_value(\n            ds,\n            \"BitsAllocated\"\n        ),\n\n        \"BitsStored\": safe_value(\n            ds,\n            \"BitsStored\"\n        ),\n\n        \"HighBit\": safe_value(\n            ds,\n            \"HighBit\"\n        ),\n\n        \"PixelRepresentation\": safe_value(\n            ds,\n            \"PixelRepresentation\"\n        ),\n\n        \"PixelSpacing\": safe_value(\n            ds,\n            \"PixelSpacing\"\n        ),\n\n        \"SliceThickness\": safe_value(\n            ds,\n            \"SliceThickness\"\n        ),\n\n        \"SpacingBetweenSlices\": safe_value(\n            ds,\n            \"SpacingBetweenSlices\"\n        ),\n\n        \"ImageOrientationPatient\": safe_value(\n            ds,\n            \"ImageOrientationPatient\"\n        ),\n\n        \"ImagePositionPatient\": safe_value(\n            ds,\n            \"ImagePositionPatient\"\n        ),\n\n        \"InstanceNumber\": safe_value(\n            ds,\n            \"InstanceNumber\"\n        ),\n\n        \"SeriesDescription\": safe_value(\n            ds,\n            \"SeriesDescription\"\n        ),\n\n        \"ProtocolName\": safe_value(\n            ds,\n            \"ProtocolName\"\n        ),\n\n        \"Manufacturer\": safe_value(\n            ds,\n            \"Manufacturer\"\n        ),\n\n        \"ManufacturerModelName\": safe_value(\n            ds,\n            \"ManufacturerModelName\"\n        ),\n\n        \"RescaleSlope\": safe_value(\n            ds,\n            \"RescaleSlope\"\n        ),\n\n        \"RescaleIntercept\": safe_value(\n            ds,\n            \"RescaleIntercept\"\n        )\n    }\n\n\n    # --------------------------------------------------------\n    # Pixel statistics\n    # --------------------------------------------------------\n\n    pixel_stats = get_pixel_statistics(\n        ds\n    )\n\n    row.update(\n        pixel_stats\n    )\n\n\n    # --------------------------------------------------------\n    # CSV metadata\n    # --------------------------------------------------------\n\n    if metadata_row is not None:\n\n        row[\"csv_Fluid_Sensitive\"] = (\n            metadata_row.get(\n                \"Fluid_Sensitive\",\n                np.nan\n            )\n        )\n\n        row[\"csv_Fat_Suppression\"] = (\n            metadata_row.get(\n                \"Fat_Suppression\",\n                np.nan\n            )\n        )\n\n        row[\"csv_Anatomical_Plane\"] = (\n            metadata_row.get(\n                \"Anatomical_Plane\",\n                None\n            )\n        )\n\n        row[\"csv_slice_count\"] = (\n            metadata_row.get(\n                \"slice_count\",\n                np.nan\n            )\n        )\n\n\n    return row\n\n\n# ------------------------------------------------------------\n# 5. Audit training sample\n# ------------------------------------------------------------\n\nprint(\"\\nAuditing sampled training series...\")\n\nfor i, (_, r) in enumerate(\n    train_sample.iterrows(),\n    start=1\n):\n\n    row = audit_series(\n        TRAIN_SERIES_DIR,\n        r[\"StudyInstanceUID\"],\n        r[\"SeriesInstanceUID\"],\n        \"train\",\n        r\n    )\n\n    audit_rows.append(\n        row\n    )\n\n    if i % 25 == 0:\n        print(\n            f\"Processed {i} / \"\n            f\"{len(train_sample)} train series\"\n        )\n\n\n# ------------------------------------------------------------\n# 6. Audit ALL test series\n# ------------------------------------------------------------\n\nprint(\"\\nAuditing ALL test series...\")\n\nfor i, (_, r) in enumerate(\n    test_series.iterrows(),\n    start=1\n):\n\n    row = audit_series(\n        TEST_SERIES_DIR,\n        r[\"StudyInstanceUID\"],\n        r[\"SeriesInstanceUID\"],\n        \"test\",\n        r\n    )\n\n    audit_rows.append(\n        row\n    )\n\n    print(\n        f\"Processed test series \"\n        f\"{i} / {len(test_series)}\"\n    )\n\n\n# ------------------------------------------------------------\n# 7. Save audit\n# ------------------------------------------------------------\n\ndicom_audit = pd.DataFrame(\n    audit_rows\n)\n\ndicom_audit.to_csv(\n    OUTPUT_FILE,\n    index=False\n)\n\n\n# ------------------------------------------------------------\n# 8. Summary\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"DICOM AUDIT SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    \"Rows audited:\",\n    len(dicom_audit)\n)\n\nprint(\n    \"Successfully read:\",\n    loaded_count\n)\n\nprint(\n    \"Failed:\",\n    failed_count\n)\n\nprint(\n    \"Missing PixelData:\",\n    missing_pixel_count\n)\n\n\nprint(\"\\nRead status:\")\n\nprint(\n    dicom_audit[\n        \"read_status\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\nprint(\"\\nDimensions:\")\n\nprint(\n    dicom_audit[\n        [\n            \"Rows\",\n            \"Columns\"\n        ]\n    ].value_counts(\n        dropna=False\n    ).head(20)\n)\n\n\nprint(\"\\nPhotometric interpretation:\")\n\nprint(\n    dicom_audit[\n        \"PhotometricInterpretation\"\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\nprint(\"\\nTransfer Syntax:\")\n\nprint(\n    dicom_audit[\n        \"TransferSyntaxUID\"\n    ].value_counts(\n        dropna=False\n    ).head(20)\n)\n\n\nprint(\"\\nPixel representation:\")\n\nprint(\n    dicom_audit[\n        [\n            \"BitsAllocated\",\n            \"BitsStored\",\n            \"PixelRepresentation\"\n        ]\n    ].value_counts(\n        dropna=False\n    )\n)\n\n\nprint(\"\\nCSV plane distribution in audited sample:\")\n\nif \"csv_Anatomical_Plane\" in dicom_audit.columns:\n\n    print(\n        dicom_audit[\n            \"csv_Anatomical_Plane\"\n        ].value_counts(\n            dropna=False\n        )\n    )\n\n\nprint(\"\\nPixel intensity summary:\")\n\nprint(\n    dicom_audit[\n        [\n            \"pixel_min\",\n            \"pixel_max\",\n            \"pixel_mean\",\n            \"pixel_std\",\n            \"pixel_p01\",\n            \"pixel_p50\",\n            \"pixel_p99\"\n        ]\n    ].describe()\n)\n\n\nprint(\"\\nMissing metadata:\")\n\nmetadata_fields = [\n    \"Rows\",\n    \"Columns\",\n    \"PixelSpacing\",\n    \"SliceThickness\",\n    \"SpacingBetweenSlices\",\n    \"ImageOrientationPatient\",\n    \"ImagePositionPatient\",\n    \"InstanceNumber\",\n    \"PhotometricInterpretation\",\n    \"BitsAllocated\",\n    \"BitsStored\",\n    \"PixelRepresentation\",\n    \"TransferSyntaxUID\",\n    \"SeriesDescription\"\n]\n\nmissing_summary = []\n\nfor field in metadata_fields:\n\n    missing = (\n        dicom_audit[field]\n        .isna()\n        |\n        (\n            dicom_audit[field]\n            .astype(str)\n            .str.strip()\n            .isin(\n                [\n                    \"\",\n                    \"None\",\n                    \"nan\"\n                ]\n            )\n        )\n    ).sum()\n\n    missing_summary.append({\n        \"field\": field,\n        \"missing_count\": int(missing),\n        \"missing_percentage\": (\n            100 * missing /\n            len(dicom_audit)\n        )\n    })\n\nmissing_df = pd.DataFrame(\n    missing_summary\n)\n\nprint(\n    missing_df.to_string(\n        index=False\n    )\n)\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 20 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"\\nSaved:\")\nprint(OUTPUT_FILE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:30:51.62638Z","iopub.execute_input":"2026-08-11T08:30:51.62676Z","iopub.status.idle":"2026-08-11T08:30:58.100067Z","shell.execute_reply.started":"2026-08-11T08:30:51.626722Z","shell.execute_reply":"2026-08-11T08:30:58.098967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 21: REPORT SEMANTIC + TERMINOLOGY AUDIT\n# ============================================================\n\nimport re\nimport unicodedata\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom collections import Counter\n\nprint(\"=\" * 70)\nprint(\"REPORT SEMANTIC + TERMINOLOGY AUDIT\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------------------\n# 1. Configuration\n# ------------------------------------------------------------\n\nTRAIN_CSV = (\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train.csv\"\n)\n\nAUDIT_DIR = Path(\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nAUDIT_DIR.mkdir(\n    parents=True,\n    exist_ok=True\n)\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\n# ------------------------------------------------------------\n# 2. Load train.csv independently\n# ------------------------------------------------------------\n\ntrain = pd.read_csv(\n    TRAIN_CSV\n)\n\nprint(\"\\nTrain shape:\", train.shape)\n\nprint(\n    \"Reports available:\",\n    train[\"Report\"].notna().sum()\n)\n\nprint(\n    \"Missing reports:\",\n    train[\"Report\"].isna().sum()\n)\n\n\n# ------------------------------------------------------------\n# 3. Basic report statistics\n# ------------------------------------------------------------\n\ntrain[\"Report\"] = (\n    train[\"Report\"]\n    .fillna(\"\")\n    .astype(str)\n)\n\ntrain[\"report_chars\"] = (\n    train[\"Report\"]\n    .str.len()\n)\n\ntrain[\"report_words\"] = (\n    train[\"Report\"]\n    .str.split()\n    .str.len()\n)\n\ntrain[\"report_lines\"] = (\n    train[\"Report\"]\n    .str.count(\"\\n\") + 1\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"REPORT LENGTH STATISTICS\")\nprint(\"=\" * 70)\n\nprint(\n    train[\n        [\n            \"report_chars\",\n            \"report_words\",\n            \"report_lines\"\n        ]\n    ].describe()\n)\n\n\n# ------------------------------------------------------------\n# 4. Script detection\n# ------------------------------------------------------------\n\ndef detect_script(text):\n\n    counts = Counter()\n\n    for char in text:\n\n        if char.isspace():\n            continue\n\n        name = unicodedata.name(\n            char,\n            \"\"\n        )\n\n        if \"LATIN\" in name:\n            counts[\"LATIN\"] += 1\n\n        elif \"GREEK\" in name:\n            counts[\"GREEK\"] += 1\n\n        elif \"CYRILLIC\" in name:\n            counts[\"CYRILLIC\"] += 1\n\n        elif \"ARABIC\" in name:\n            counts[\"ARABIC\"] += 1\n\n        elif \"HEBREW\" in name:\n            counts[\"HEBREW\"] += 1\n\n        elif \"DEVANAGARI\" in name:\n            counts[\"DEVANAGARI\"] += 1\n\n        else:\n            counts[\"OTHER\"] += 1\n\n    if not counts:\n        return \"EMPTY\"\n\n    return counts.most_common(1)[0][0]\n\n\ntrain[\"dominant_script\"] = (\n    train[\"Report\"]\n    .apply(detect_script)\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"DOMINANT REPORT SCRIPT\")\nprint(\"=\" * 70)\n\nprint(\n    train[\"dominant_script\"]\n    .value_counts()\n)\n\n\n# ------------------------------------------------------------\n# 5. Target terminology\n# ------------------------------------------------------------\n\n# These are deliberately broad candidate terms.\n# They are NOT labels.\n# Cell 22 will validate them against the 58 labeled studies.\n\nTERM_PATTERNS = {\n\n    \"ACL\": [\n        r\"\\bacl\\b\",\n        r\"anterior cruciate\",\n        r\"ligament croisé antérieur\",\n        r\"ligamento cruzado anterior\",\n        r\"vorderes kreuzband\"\n    ],\n\n    \"MCL\": [\n        r\"\\bmcl\\b\",\n        r\"medial collateral\",\n        r\"medial collateral ligament\",\n        r\"ligament collatéral médial\",\n        r\"ligamento colateral medial\"\n    ],\n\n    \"Medial Meniscus\": [\n        r\"medial meniscus\",\n        r\"medial meniscal\",\n        r\"menisque médial\",\n        r\"menisco medial\",\n        r\"innenmeniskus\"\n    ],\n\n    \"Lateral Meniscus\": [\n        r\"lateral meniscus\",\n        r\"lateral meniscal\",\n        r\"menisque latéral\",\n        r\"menisco lateral\",\n        r\"außenmeniskus\"\n    ],\n\n    \"Medial OA\": [\n        r\"medial compartment\",\n        r\"medial femorotibial\",\n        r\"medial osteoarthritis\",\n        r\"medial arthrosis\",\n        r\"medial gonarthrosis\",\n        r\"medial joint space narrowing\"\n    ],\n\n    \"Lateral OA\": [\n        r\"lateral compartment\",\n        r\"lateral femorotibial\",\n        r\"lateral osteoarthritis\",\n        r\"lateral arthrosis\",\n        r\"lateral gonarthrosis\",\n        r\"lateral joint space narrowing\"\n    ],\n\n    \"PF OA\": [\n        r\"patellofemoral\",\n        r\"patello-femoral\",\n        r\"patellofemoral osteoarthritis\",\n        r\"patellofemoral arthrosis\",\n        r\"retropatellar\",\n        r\"patellar cartilage\"\n    ],\n\n    \"Effusion\": [\n        r\"\\beffusion\\b\",\n        r\"joint effusion\",\n        r\"articular effusion\",\n        r\"épanchement\",\n        r\"gelenkerguss\"\n    ],\n\n    \"Synovitis\": [\n        r\"\\bsynovitis\\b\",\n        r\"synovial\",\n        r\"synoviale\",\n        r\"synoviale verdickung\"\n    ],\n\n    \"Baker's\": [\n        r\"baker\",\n        r\"popliteal cyst\",\n        r\"popliteal cyste\",\n        r\"cyste poplitée\",\n        r\"baker cyst\"\n    ],\n\n    \"Contusion\": [\n        r\"\\bcontusion\\b\",\n        r\"bone bruise\",\n        r\"bone marrow edema\",\n        r\"marrow edema\",\n        r\"œdème osseux\",\n        r\"knochenödem\"\n    ],\n\n    \"Fracture\": [\n        r\"\\bfracture\\b\",\n        r\"\\bfractures\\b\",\n        r\"fracture line\",\n        r\"fracture du\",\n        r\"fraktur\",\n        r\"fractura\"\n    ]\n}\n\n\n# ------------------------------------------------------------\n# 6. Compile patterns\n# ------------------------------------------------------------\n\ncompiled_patterns = {}\n\nfor target, patterns in TERM_PATTERNS.items():\n\n    compiled_patterns[target] = [\n        re.compile(\n            pattern,\n            flags=re.IGNORECASE\n        )\n        for pattern in patterns\n    ]\n\n\n# ------------------------------------------------------------\n# 7. Negation / uncertainty patterns\n# ------------------------------------------------------------\n\nNEGATION_PATTERNS = [\n\n    r\"\\bno\\b\",\n    r\"\\bwithout\\b\",\n    r\"\\bnot\\b\",\n    r\"\\bnegative\\b\",\n    r\"\\bnormal\\b\",\n    r\"\\bintact\\b\",\n    r\"\\bpreserved\\b\",\n    r\"\\babsence\\b\",\n    r\"\\babsent\\b\",\n\n    r\"\\baucune\\b\",\n    r\"\\baucun\\b\",\n    r\"\\bsans\\b\",\n    r\"\\bnormal\\b\",\n    r\"\\bintact\\b\",\n\n    r\"\\bkein\\b\",\n    r\"\\bkeine\\b\",\n    r\"\\bkeinen\\b\",\n    r\"\\bintakt\\b\",\n\n    r\"\\bno hay\\b\",\n    r\"\\bsin\\b\",\n    r\"\\bnormal\\b\",\n    r\"\\bíntegro\\b\",\n    r\"\\bintegro\\b\"\n]\n\nUNCERTAINTY_PATTERNS = [\n\n    r\"\\bpossible\\b\",\n    r\"\\bpossibly\\b\",\n    r\"\\bprobable\\b\",\n    r\"\\bprobable\\b\",\n    r\"\\bsuspect\\b\",\n    r\"\\bsuspicious\\b\",\n    r\"\\bcannot exclude\\b\",\n    r\"\\bmay represent\\b\",\n    r\"\\bmay be\\b\",\n\n    r\"\\bpossible\\b\",\n    r\"\\bsuspect\\b\",\n    r\"\\bà confirmer\\b\",\n    r\"\\bne peut être exclu\\b\",\n\n    r\"\\bmöglich\\b\",\n    r\"\\bverdacht\\b\",\n\n    r\"\\bposible\\b\",\n    r\"\\bprobable\\b\",\n    r\"\\bno se puede excluir\\b\"\n]\n\ncompiled_negation = [\n    re.compile(\n        p,\n        flags=re.IGNORECASE\n    )\n    for p in NEGATION_PATTERNS\n]\n\ncompiled_uncertainty = [\n    re.compile(\n        p,\n        flags=re.IGNORECASE\n    )\n    for p in UNCERTAINTY_PATTERNS\n]\n\n\n# ------------------------------------------------------------\n# 8. Find target terminology\n# ------------------------------------------------------------\n\ndef find_matches(text, patterns):\n\n    matches = []\n\n    for pattern in patterns:\n\n        for match in pattern.finditer(text):\n\n            start = max(\n                0,\n                match.start() - 80\n            )\n\n            end = min(\n                len(text),\n                match.end() + 80\n            )\n\n            context = text[\n                start:end\n            ].replace(\n                \"\\n\",\n                \" \"\n            )\n\n            matches.append({\n                \"term\": match.group(),\n                \"context\": context\n            })\n\n    return matches\n\n\n# ------------------------------------------------------------\n# 9. Build terminology audit\n# ------------------------------------------------------------\n\naudit_rows = []\n\nfor _, row in train.iterrows():\n\n    text = row[\"Report\"]\n\n    for target in TARGETS:\n\n        matches = find_matches(\n            text,\n            compiled_patterns[target]\n        )\n\n        if matches:\n\n            context_text = \" || \".join(\n                m[\"context\"]\n                for m in matches[:5]\n            )\n\n            negation_found = any(\n                pattern.search(\n                    context_text\n                )\n                for pattern in compiled_negation\n            )\n\n            uncertainty_found = any(\n                pattern.search(\n                    context_text\n                )\n                for pattern in compiled_uncertainty\n            )\n\n            audit_rows.append({\n\n                \"StudyInstanceUID\":\n                    row[\"StudyInstanceUID\"],\n\n                \"target\":\n                    target,\n\n                \"mention_count\":\n                    len(matches),\n\n                \"matched_terms\":\n                    \" | \".join(\n                        sorted(\n                            set(\n                                m[\"term\"]\n                                for m in matches\n                            )\n                        )\n                    ),\n\n                \"negation_context_found\":\n                    negation_found,\n\n                \"uncertainty_context_found\":\n                    uncertainty_found,\n\n                \"context\":\n                    context_text\n            })\n\n\nreport_term_audit = pd.DataFrame(\n    audit_rows\n)\n\n\n# ------------------------------------------------------------\n# 10. Summary by target\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TARGET TERMINOLOGY SUMMARY\")\nprint(\"=\" * 70)\n\nif len(report_term_audit) > 0:\n\n    target_summary = (\n        report_term_audit\n        .groupby(\"target\")\n        .agg(\n            studies_with_mentions=(\n                \"StudyInstanceUID\",\n                \"nunique\"\n            ),\n            total_mentions=(\n                \"mention_count\",\n                \"sum\"\n            ),\n            negation_contexts=(\n                \"negation_context_found\",\n                \"sum\"\n            ),\n            uncertainty_contexts=(\n                \"uncertainty_context_found\",\n                \"sum\"\n            )\n        )\n        .reindex(TARGETS)\n        .fillna(0)\n    )\n\n    print(\n        target_summary\n    )\n\nelse:\n\n    target_summary = pd.DataFrame()\n\n\n# ------------------------------------------------------------\n# 11. Explicitly labeled subset\n# ------------------------------------------------------------\n\nlabel_mask = train[\n    TARGETS\n].notna().all(axis=1)\n\nlabeled_train = train.loc[\n    label_mask\n].copy()\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"EXPLICITLY LABELED REPORT AUDIT\")\nprint(\"=\" * 70)\n\nprint(\n    \"Explicitly labeled studies:\",\n    len(labeled_train)\n)\n\nprint(\n    \"Labeled studies with reports:\",\n    labeled_train[\"Report\"]\n    .notna()\n    .sum()\n)\n\n\n# ------------------------------------------------------------\n# 12. Display representative reports\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"REPRESENTATIVE REPORTS FROM LABELED STUDIES\")\nprint(\"=\" * 70)\n\ndisplay_columns = [\n    \"StudyInstanceUID\",\n    \"Report\"\n] + TARGETS\n\nsample_labeled = labeled_train.sample(\n    n=min(10, len(labeled_train)),\n    random_state=42\n)\n\nfor i, (_, row) in enumerate(\n    sample_labeled.iterrows(),\n    start=1\n):\n\n    print(\"\\n\" + \"-\" * 70)\n    print(\n        f\"Labeled report {i}\"\n    )\n    print(\n        \"StudyInstanceUID:\",\n        row[\"StudyInstanceUID\"]\n    )\n\n    print(\"\\nLabels:\")\n\n    for target in TARGETS:\n\n        print(\n            f\"  {target}:\",\n            int(row[target])\n        )\n\n    print(\"\\nReport:\")\n    print(\n        row[\"Report\"]\n    )\n\n\n# ------------------------------------------------------------\n# 13. Save outputs\n# ------------------------------------------------------------\n\nterm_file = (\n    AUDIT_DIR /\n    \"report_terminology_audit.csv\"\n)\n\nsummary_file = (\n    AUDIT_DIR /\n    \"report_terminology_summary.csv\"\n)\n\nreport_term_audit.to_csv(\n    term_file,\n    index=False\n)\n\ntarget_summary.to_csv(\n    summary_file\n)\n\n\n# ------------------------------------------------------------\n# 14. Final output\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 21 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"\\nSaved:\")\nprint(term_file)\nprint(summary_file)\n\nprint(\"\\nTotal report-target mention records:\")\nprint(len(report_term_audit))\n\nprint(\"\\nImportant:\")\nprint(\n    \"These terminology matches are CANDIDATE EVIDENCE only.\"\n)\n\nprint(\n    \"They must NOT be treated as ground-truth labels yet.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:33:25.955998Z","iopub.execute_input":"2026-08-11T08:33:25.956335Z","iopub.status.idle":"2026-08-11T08:33:34.59629Z","shell.execute_reply.started":"2026-08-11T08:33:25.956306Z","shell.execute_reply":"2026-08-11T08:33:34.595402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 22: REPORT ↔ LABEL CONSISTENCY AUDIT\n# ============================================================\n\nimport re\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\nprint(\"=\" * 70)\nprint(\"REPORT ↔ LABEL CONSISTENCY AUDIT\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------------------\n# 1. Configuration\n# ------------------------------------------------------------\n\nTRAIN_CSV = (\n    \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train.csv\"\n)\n\nAUDIT_DIR = Path(\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\n# ------------------------------------------------------------\n# 2. Load data\n# ------------------------------------------------------------\n\ntrain = pd.read_csv(\n    TRAIN_CSV\n)\n\ntrain[\"Report\"] = (\n    train[\"Report\"]\n    .fillna(\"\")\n    .astype(str)\n)\n\n# Explicitly labeled studies\nlabeled = train[\n    train[TARGETS].notna().all(axis=1)\n].copy()\n\nprint(\n    \"\\nExplicitly labeled studies:\",\n    len(labeled)\n)\n\n# ------------------------------------------------------------\n# 3. Target-specific positive evidence\n# ------------------------------------------------------------\n\nPOSITIVE_PATTERNS = {\n\n    \"ACL\": [\n        r\"acl.*tear\",\n        r\"tear.*acl\",\n        r\"acl.*ruptur\",\n        r\"ruptur.*acl\",\n        r\"acl.*injur\",\n        r\"acl.*sprain\",\n        r\"anterior cruciate ligament.*tear\",\n        r\"anterior cruciate ligament.*ruptur\",\n        r\"anterior cruciate ligament.*injur\",\n        r\"ligamento cruzado anterior.*rotura\",\n        r\"ligament.*croisé antérieur.*rupture\"\n    ],\n\n    \"MCL\": [\n        r\"mcl.*tear\",\n        r\"mcl.*injur\",\n        r\"mcl.*sprain\",\n        r\"medial collateral ligament.*tear\",\n        r\"medial collateral ligament.*injur\",\n        r\"medial collateral ligament.*sprain\",\n        r\"ligamento colateral medial.*esguince\",\n        r\"ligamento colateral medial.*rotura\",\n        r\"ligamento colateral medial.*lesion\",\n        r\"ligament.*collatéral médial.*lésion\"\n    ],\n\n    \"Medial Meniscus\": [\n        r\"medial meniscus.*tear\",\n        r\"medial meniscal.*tear\",\n        r\"medial meniscus.*ruptur\",\n        r\"medial meniscus.*lesion\",\n        r\"meniscus.*medial.*tear\",\n        r\"menisco medial.*rotura\",\n        r\"menisco medial.*desgarro\",\n        r\"menisque médial.*rupture\",\n        r\"innenmeniskus.*riss\"\n    ],\n\n    \"Lateral Meniscus\": [\n        r\"lateral meniscus.*tear\",\n        r\"lateral meniscal.*tear\",\n        r\"lateral meniscus.*ruptur\",\n        r\"lateral meniscus.*lesion\",\n        r\"meniscus.*lateral.*tear\",\n        r\"menisco lateral.*rotura\",\n        r\"menisco lateral.*desgarro\",\n        r\"menisque latéral.*rupture\",\n        r\"außenmeniskus.*riss\"\n    ],\n\n    \"Medial OA\": [\n        r\"medial compartment.*osteoarthritis\",\n        r\"medial compartment.*arthrosis\",\n        r\"medial compartment.*chondrosis\",\n        r\"medial compartment.*chondral\",\n        r\"medial.*joint space narrowing\",\n        r\"medial.*cartilage loss\",\n        r\"medial.*osteophyte\",\n        r\"medial tibiofemoral.*osteoarthritis\",\n        r\"medial femorotibial.*osteoarthritis\",\n        r\"osteoarthritis.*medial compartment\",\n        r\"arthrosis.*medial compartment\"\n    ],\n\n    \"Lateral OA\": [\n        r\"lateral compartment.*osteoarthritis\",\n        r\"lateral compartment.*arthrosis\",\n        r\"lateral compartment.*chondrosis\",\n        r\"lateral compartment.*chondral\",\n        r\"lateral.*joint space narrowing\",\n        r\"lateral.*cartilage loss\",\n        r\"lateral.*osteophyte\",\n        r\"lateral tibiofemoral.*osteoarthritis\",\n        r\"lateral femorotibial.*osteoarthritis\",\n        r\"osteoarthritis.*lateral compartment\",\n        r\"arthrosis.*lateral compartment\"\n    ],\n\n    \"PF OA\": [\n        r\"patellofemoral.*osteoarthritis\",\n        r\"patellofemoral.*arthrosis\",\n        r\"patellofemoral.*chondrosis\",\n        r\"patellofemoral.*chondral\",\n        r\"patellofemoral.*cartilage loss\",\n        r\"patellofemoral.*cartilage.*damage\",\n        r\"patellar.*cartilage loss\",\n        r\"patellar.*chondrosis\",\n        r\"retropatellar.*chondrosis\",\n        r\"patellofemoral compartment.*degener\"\n    ],\n\n    \"Effusion\": [\n        r\"\\beffusion\\b\",\n        r\"joint effusion\",\n        r\"knee effusion\",\n        r\"articular effusion\",\n        r\"joint fluid\",\n        r\"fluid.*joint\",\n        r\"derrame articular\",\n        r\"derrame de la articulación\",\n        r\"épanchement articulaire\",\n        r\"gelenkerguss\"\n    ],\n\n    \"Synovitis\": [\n        r\"\\bsynovitis\\b\",\n        r\"synovial.*thickening\",\n        r\"synovial.*proliferation\",\n        r\"synoviale.*verdick\",\n        r\"sinovitis\"\n    ],\n\n    \"Baker's\": [\n        r\"baker.?s cyst\",\n        r\"baker cyst\",\n        r\"popliteal cyst\",\n        r\"popliteal cyste\",\n        r\"cyste poplitée\",\n        r\"quiste poplíteo\",\n        r\"baker.*kyst\"\n    ],\n\n    \"Contusion\": [\n        r\"bone contusion\",\n        r\"bone bruise\",\n        r\"bone marrow edema\",\n        r\"marrow edema\",\n        r\"osteochondral impaction injury\",\n        r\"contusion\",\n        r\"contusion.*bone\",\n        r\"bone.*contusion\",\n        r\"œdème osseux\",\n        r\"edema.*médula ósea\"\n    ],\n\n    \"Fracture\": [\n        r\"\\bfracture\\b\",\n        r\"\\bfractures\\b\",\n        r\"fracture line\",\n        r\"fracture.*bone\",\n        r\"insufficiency fracture\",\n        r\"stress fracture\",\n        r\"microfracture\",\n        r\"microtrabecular fracture\",\n        r\"fraktur\",\n        r\"fractura\"\n    ]\n}\n\n\n# ------------------------------------------------------------\n# 4. Target-specific negative evidence\n# ------------------------------------------------------------\n\nNEGATIVE_PATTERNS = {\n\n    \"ACL\": [\n        r\"acl.*normal\",\n        r\"acl.*intact\",\n        r\"acl.*preserved\",\n        r\"acl.*no tear\",\n        r\"acl.*without tear\",\n        r\"no.*acl.*tear\",\n        r\"anterior cruciate ligament.*normal\",\n        r\"anterior cruciate ligament.*intact\",\n        r\"anterior cruciate ligament.*no tear\"\n    ],\n\n    \"MCL\": [\n        r\"mcl.*normal\",\n        r\"mcl.*intact\",\n        r\"mcl.*preserved\",\n        r\"mcl.*no tear\",\n        r\"no.*mcl.*tear\",\n        r\"medial collateral ligament.*normal\",\n        r\"medial collateral ligament.*intact\",\n        r\"medial collateral ligament.*no tear\"\n    ],\n\n    \"Medial Meniscus\": [\n        r\"medial meniscus.*normal\",\n        r\"medial meniscus.*intact\",\n        r\"medial meniscus.*no tear\",\n        r\"medial meniscus.*not torn\",\n        r\"medial meniscus.*without tear\",\n        r\"meniscus.*medial.*normal\",\n        r\"menisco medial.*normal\",\n        r\"menisco medial.*sin.*rotura\"\n    ],\n\n    \"Lateral Meniscus\": [\n        r\"lateral meniscus.*normal\",\n        r\"lateral meniscus.*intact\",\n        r\"lateral meniscus.*no tear\",\n        r\"lateral meniscus.*not torn\",\n        r\"lateral meniscus.*without tear\",\n        r\"meniscus.*lateral.*normal\",\n        r\"menisco lateral.*normal\"\n    ],\n\n    \"Medial OA\": [\n        r\"medial compartment.*normal\",\n        r\"medial compartment.*intact\",\n        r\"medial compartment.*no.*chondrosis\",\n        r\"medial compartment.*no.*chondral\",\n        r\"medial compartment.*cartilage.*intact\",\n        r\"medial compartment.*without.*chondral\"\n    ],\n\n    \"Lateral OA\": [\n        r\"lateral compartment.*normal\",\n        r\"lateral compartment.*intact\",\n        r\"lateral compartment.*no.*chondrosis\",\n        r\"lateral compartment.*no.*chondral\",\n        r\"lateral compartment.*cartilage.*intact\",\n        r\"lateral compartment.*without.*chondral\"\n    ],\n\n    \"PF OA\": [\n        r\"patellofemoral.*normal\",\n        r\"patellofemoral.*intact\",\n        r\"patellofemoral.*without.*chondral\",\n        r\"patellofemoral.*no.*chondral\",\n        r\"patellar.*cartilage.*normal\",\n        r\"patellar.*cartilage.*intact\"\n    ],\n\n    \"Effusion\": [\n        r\"no joint effusion\",\n        r\"no knee effusion\",\n        r\"no articular effusion\",\n        r\"without.*effusion\",\n        r\"no.*joint fluid\",\n        r\"no.*fluid.*joint\",\n        r\"no hay.*derrame\",\n        r\"sin.*derrame\",\n        r\"no.*épanchement\"\n    ],\n\n    \"Synovitis\": [\n        r\"no synovitis\",\n        r\"without synovitis\",\n        r\"no.*synovial.*thickening\",\n        r\"sinovitis.*no\",\n        r\"no evidence.*synovitis\"\n    ],\n\n    \"Baker's\": [\n        r\"no baker.?s cyst\",\n        r\"no popliteal cyst\",\n        r\"without.*baker.?s cyst\",\n        r\"no.*popliteal cyst\",\n        r\"no hay.*quiste poplíteo\",\n        r\"sin.*quiste poplíteo\"\n    ],\n\n    \"Contusion\": [\n        r\"no bone contusion\",\n        r\"no bone bruise\",\n        r\"no marrow edema\",\n        r\"without.*bone contusion\",\n        r\"without.*bone bruise\",\n        r\"no.*osseous.*contusion\",\n        r\"no hay.*contusion\"\n    ],\n\n    \"Fracture\": [\n        r\"no fracture\",\n        r\"no fractures\",\n        r\"without fracture\",\n        r\"without.*fracture\",\n        r\"no acute fracture\",\n        r\"no evidence.*fracture\",\n        r\"sin fractura\",\n        r\"sin.*fractura\",\n        r\"keine fraktur\"\n    ]\n}\n\n\n# ------------------------------------------------------------\n# 5. Compile regex patterns\n# ------------------------------------------------------------\n\ncompiled_positive = {\n    target: [\n        re.compile(\n            pattern,\n            flags=re.IGNORECASE\n        )\n        for pattern in patterns\n    ]\n    for target, patterns in POSITIVE_PATTERNS.items()\n}\n\ncompiled_negative = {\n    target: [\n        re.compile(\n            pattern,\n            flags=re.IGNORECASE\n        )\n        for pattern in patterns\n    ]\n    for target, patterns in NEGATIVE_PATTERNS.items()\n}\n\n\n# ------------------------------------------------------------\n# 6. Sentence splitting\n# ------------------------------------------------------------\n\ndef split_sentences(text):\n\n    text = re.sub(\n        r\"\\s+\",\n        \" \",\n        text\n    ).strip()\n\n    if not text:\n        return []\n\n    return [\n        sentence.strip()\n        for sentence in re.split(\n            r\"(?<=[.!?])\\s+\",\n            text\n        )\n        if sentence.strip()\n    ]\n\n\n# ------------------------------------------------------------\n# 7. Detect evidence\n# ------------------------------------------------------------\n\ndef detect_evidence(\n    text,\n    target\n):\n\n    sentences = split_sentences(\n        text\n    )\n\n    positive_hits = []\n    negative_hits = []\n\n    for sentence in sentences:\n\n        for pattern in compiled_positive[target]:\n\n            if pattern.search(sentence):\n\n                positive_hits.append(\n                    sentence\n                )\n                break\n\n        for pattern in compiled_negative[target]:\n\n            if pattern.search(sentence):\n\n                negative_hits.append(\n                    sentence\n                )\n                break\n\n    return (\n        positive_hits,\n        negative_hits\n    )\n\n\n# ------------------------------------------------------------\n# 8. Audit all 58 labeled studies\n# ------------------------------------------------------------\n\naudit_rows = []\n\nfor _, row in labeled.iterrows():\n\n    text = row[\"Report\"]\n\n    for target in TARGETS:\n\n        label = int(\n            row[target]\n        )\n\n        positive_hits, negative_hits = (\n            detect_evidence(\n                text,\n                target\n            )\n        )\n\n        has_positive = (\n            len(positive_hits) > 0\n        )\n\n        has_negative = (\n            len(negative_hits) > 0\n        )\n\n        if has_positive and has_negative:\n\n            evidence_status = \"MIXED\"\n\n        elif has_positive:\n\n            evidence_status = \"POSITIVE_EVIDENCE\"\n\n        elif has_negative:\n\n            evidence_status = \"NEGATIVE_EVIDENCE\"\n\n        else:\n\n            evidence_status = \"NO_CLEAR_EVIDENCE\"\n\n\n        # ----------------------------------------------------\n        # Compare evidence with actual label\n        # ----------------------------------------------------\n\n        if label == 1:\n\n            if has_positive:\n                consistency = \"AGREEMENT_POSITIVE\"\n\n            elif has_negative:\n                consistency = \"CONFLICT_LABEL_1_REPORT_NEGATIVE\"\n\n            else:\n                consistency = \"UNCLEAR_LABEL_1\"\n\n        else:\n\n            if has_positive:\n                consistency = \"CONFLICT_LABEL_0_REPORT_POSITIVE\"\n\n            elif has_negative:\n                consistency = \"AGREEMENT_NEGATIVE\"\n\n            else:\n                consistency = \"UNCLEAR_LABEL_0\"\n\n\n        audit_rows.append({\n\n            \"StudyInstanceUID\":\n                row[\"StudyInstanceUID\"],\n\n            \"target\":\n                target,\n\n            \"label\":\n                label,\n\n            \"positive_evidence\":\n                has_positive,\n\n            \"negative_evidence\":\n                has_negative,\n\n            \"evidence_status\":\n                evidence_status,\n\n            \"consistency\":\n                consistency,\n\n            \"positive_context\":\n                \" || \".join(\n                    positive_hits[:3]\n                ),\n\n            \"negative_context\":\n                \" || \".join(\n                    negative_hits[:3]\n                )\n        })\n\n\nconsistency_df = pd.DataFrame(\n    audit_rows\n)\n\n\n# ------------------------------------------------------------\n# 9. Overall consistency summary\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OVERALL CONSISTENCY SUMMARY\")\nprint(\"=\" * 70)\n\noverall_summary = (\n    consistency_df[\n        \"consistency\"\n    ]\n    .value_counts()\n)\n\nprint(\n    overall_summary\n)\n\n\n# ------------------------------------------------------------\n# 10. Consistency by target\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CONSISTENCY BY TARGET\")\nprint(\"=\" * 70)\n\ntarget_consistency = pd.crosstab(\n    consistency_df[\"target\"],\n    consistency_df[\"consistency\"]\n)\n\ntarget_consistency = (\n    target_consistency\n    .reindex(TARGETS)\n    .fillna(0)\n    .astype(int)\n)\n\nprint(\n    target_consistency\n)\n\n\n# ------------------------------------------------------------\n# 11. Calculate agreement metrics\n# ------------------------------------------------------------\n\nmetrics = []\n\nfor target in TARGETS:\n\n    subset = consistency_df[\n        consistency_df[\"target\"] == target\n    ]\n\n    label_1 = subset[\n        subset[\"label\"] == 1\n    ]\n\n    label_0 = subset[\n        subset[\"label\"] == 0\n    ]\n\n    tp_like = (\n        label_1[\"positive_evidence\"]\n        .sum()\n    )\n\n    fn_like = (\n        (~label_1[\"positive_evidence\"])\n        .sum()\n    )\n\n    fp_like = (\n        label_0[\"positive_evidence\"]\n        .sum()\n    )\n\n    tn_like = (\n        (~label_0[\"positive_evidence\"])\n        .sum()\n    )\n\n    total = len(subset)\n\n    evidence_agreement = (\n        (\n            tp_like +\n            tn_like\n        ) / total\n        if total > 0\n        else np.nan\n    )\n\n    metrics.append({\n\n        \"target\":\n            target,\n\n        \"n\":\n            total,\n\n        \"label_positive\":\n            len(label_1),\n\n        \"label_negative\":\n            len(label_0),\n\n        \"positive_evidence_when_label_1\":\n            int(tp_like),\n\n        \"positive_evidence_when_label_0\":\n            int(fp_like),\n\n        \"negative_or_no_positive_evidence_when_label_0\":\n            int(tn_like),\n\n        \"label_1_without_positive_evidence\":\n            int(fn_like),\n\n        \"positive_evidence_agreement\":\n            evidence_agreement\n    })\n\n\nmetrics_df = pd.DataFrame(\n    metrics\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"REPORT EVIDENCE AGREEMENT METRICS\")\nprint(\"=\" * 70)\n\nprint(\n    metrics_df.to_string(\n        index=False\n    )\n)\n\n\n# ------------------------------------------------------------\n# 12. Show conflicts\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LABEL 0 / REPORT POSITIVE CONFLICTS\")\nprint(\"=\" * 70)\n\nconflict_0 = consistency_df[\n    consistency_df[\"consistency\"]\n    == \"CONFLICT_LABEL_0_REPORT_POSITIVE\"\n]\n\nprint(\n    \"Count:\",\n    len(conflict_0)\n)\n\nif len(conflict_0) > 0:\n\n    display(\n        conflict_0[\n            [\n                \"StudyInstanceUID\",\n                \"target\",\n                \"label\",\n                \"positive_context\"\n            ]\n        ].head(30)\n    )\n\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LABEL 1 / REPORT NEGATIVE CONFLICTS\")\nprint(\"=\" * 70)\n\nconflict_1 = consistency_df[\n    consistency_df[\"consistency\"]\n    == \"CONFLICT_LABEL_1_REPORT_NEGATIVE\"\n]\n\nprint(\n    \"Count:\",\n    len(conflict_1)\n)\n\nif len(conflict_1) > 0:\n\n    display(\n        conflict_1[\n            [\n                \"StudyInstanceUID\",\n                \"target\",\n                \"label\",\n                \"negative_context\"\n            ]\n        ].head(30)\n    )\n\n\n# ------------------------------------------------------------\n# 13. Save outputs\n# ------------------------------------------------------------\n\nconsistency_file = (\n    AUDIT_DIR /\n    \"report_label_consistency.csv\"\n)\n\nmetrics_file = (\n    AUDIT_DIR /\n    \"report_label_consistency_summary.csv\"\n)\n\nconflict_file = (\n    AUDIT_DIR /\n    \"report_label_conflicts.csv\"\n)\n\nconsistency_df.to_csv(\n    consistency_file,\n    index=False\n)\n\nmetrics_df.to_csv(\n    metrics_file,\n    index=False\n)\n\npd.concat(\n    [\n        conflict_0,\n        conflict_1\n    ],\n    ignore_index=True\n).to_csv(\n    conflict_file,\n    index=False\n)\n\n\n# ------------------------------------------------------------\n# 14. Final status\n# ------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 22 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"\\nSaved:\")\nprint(consistency_file)\nprint(metrics_file)\nprint(conflict_file)\n\nprint(\"\\nIMPORTANT:\")\nprint(\n    \"This audit measures report-label agreement.\"\n)\n\nprint(\n    \"It does NOT create pseudo-labels.\"\n)\n\nprint(\n    \"Do not train on report-derived labels yet.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:36:10.929556Z","iopub.execute_input":"2026-08-11T08:36:10.929968Z","iopub.status.idle":"2026-08-11T08:36:11.444912Z","shell.execute_reply.started":"2026-08-11T08:36:10.929937Z","shell.execute_reply":"2026-08-11T08:36:11.444023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 23 - IMPROVED REPORT EVIDENCE AUDIT\n# ================================================================\n\nimport os\nimport re\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"IMPROVED REPORT EVIDENCE AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# 1. Make sure train and target definitions exist\n# ----------------------------------------------------------------\n\nif \"train\" not in globals():\n    train_csv_path = \"/kaggle/input/competitions/rsna-knee-abnormality-detection/train.csv\"\n\n    if not os.path.exists(train_csv_path):\n        raise FileNotFoundError(\n            f\"Could not find train.csv at:\\n{train_csv_path}\"\n        )\n\n    train = pd.read_csv(train_csv_path)\n\nif \"TARGETS\" not in globals():\n    TARGETS = [\n        \"ACL\",\n        \"MCL\",\n        \"Medial Meniscus\",\n        \"Lateral Meniscus\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Synovitis\",\n        \"Baker's\",\n        \"Contusion\",\n        \"Fracture\"\n    ]\n\nprint(f\"Train studies: {len(train)}\")\nprint(f\"Explicitly labeled studies: {train[TARGETS].notna().all(axis=1).sum()}\")\n\n\n# ----------------------------------------------------------------\n# 2. Conservative target aliases\n# ----------------------------------------------------------------\n\nTARGET_ALIASES = {\n    \"ACL\": [\n        r\"\\bacl\\b\",\n        r\"anterior cruciate ligament\"\n    ],\n\n    \"MCL\": [\n        r\"\\bmcl\\b\",\n        r\"medial collateral ligament\"\n    ],\n\n    \"Medial Meniscus\": [\n        r\"medial meniscus\",\n        r\"meniscus medialis\"\n    ],\n\n    \"Lateral Meniscus\": [\n        r\"lateral meniscus\",\n        r\"meniscus lateralis\"\n    ],\n\n    \"Medial OA\": [\n        r\"medial osteoarthritis\",\n        r\"medial compartment.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"medial compartment.*(?:joint space narrowing|osteophyte)\"\n    ],\n\n    \"Lateral OA\": [\n        r\"lateral osteoarthritis\",\n        r\"lateral compartment.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"lateral compartment.*(?:joint space narrowing|osteophyte)\"\n    ],\n\n    \"PF OA\": [\n        r\"patellofemoral osteoarthritis\",\n        r\"patellofemoral.*(?:osteoarthritis|arthrosis|degenerative)\",\n        r\"patellofemoral.*(?:joint space narrowing|osteophyte)\"\n    ],\n\n    \"Effusion\": [\n        r\"\\beffusion\\b\",\n        r\"joint effusion\"\n    ],\n\n    \"Synovitis\": [\n        r\"\\bsynovitis\\b\",\n        r\"synovial thickening\"\n    ],\n\n    \"Baker's\": [\n        r\"baker.?s cyst\",\n        r\"popliteal cyst\"\n    ],\n\n    \"Contusion\": [\n        r\"\\bcontusion\\b\",\n        r\"bone bruise\",\n        r\"bone contusion\"\n    ],\n\n    \"Fracture\": [\n        r\"\\bfracture\\b\",\n        r\"fractured\"\n    ]\n}\n\n\n# ----------------------------------------------------------------\n# 3. Explicit negative language\n# ----------------------------------------------------------------\n\nNEGATIVE_PATTERNS = [\n    r\"\\bno\\b\",\n    r\"\\bnot\\b\",\n    r\"\\bwithout\\b\",\n    r\"\\bnegative for\\b\",\n    r\"\\bno evidence of\\b\",\n    r\"\\bno sign of\\b\",\n    r\"\\bno signs of\\b\",\n    r\"\\bno acute\\b\",\n    r\"\\bintact\\b\",\n    r\"\\bnormal\\b\",\n    r\"\\bunremarkable\\b\",\n    r\"\\bpreserved\\b\",\n    r\"\\bconserved\\b\",\n    r\"\\bconservation\\b\",\n    r\"\\babsence of\\b\",\n    r\"\\babsent\\b\",\n    r\"\\bnot seen\\b\",\n    r\"\\bnone\\b\"\n]\n\n\n# ----------------------------------------------------------------\n# 4. Explicit positive/pathological language\n# ----------------------------------------------------------------\n\nPOSITIVE_PATTERNS = [\n    r\"\\btear\\b\",\n    r\"\\btear(?:ed|ing)?\\b\",\n    r\"\\bru(p|pture|ptured)\\b\",\n    r\"\\blésion\\b\",\n    r\"\\blesion\\b\",\n    r\"\\babnormal\\b\",\n    r\"\\bpathologic\\b\",\n    r\"\\bpathological\\b\",\n    r\"\\bdegenerative\\b\",\n    r\"\\bdegeneration\\b\",\n    r\"\\bosteoarthritis\\b\",\n    r\"\\barthrosis\\b\",\n    r\"\\bosteoarthritic\\b\",\n    r\"\\bswelling\\b\",\n    r\"\\bfluid\\b\",\n    r\"\\bthickening\\b\",\n    r\"\\bcontusion\\b\",\n    r\"\\bbruise\\b\",\n    r\"\\bfracture\\b\",\n    r\"\\bcyst\\b\",\n    r\"\\beffusion\\b\",\n    r\"\\bsynovitis\\b\",\n    r\"\\bsignal abnormality\\b\",\n    r\"\\bhyperintense\\b\",\n    r\"\\bdisruption\\b\",\n    r\"\\bdisrupted\\b\"\n]\n\n\n# ----------------------------------------------------------------\n# 5. Sentence splitter\n# ----------------------------------------------------------------\n\ndef split_sentences(text):\n    if pd.isna(text):\n        return []\n\n    text = str(text)\n\n    # Normalize line breaks\n    text = re.sub(r\"[\\r\\n]+\", \" \", text)\n\n    # Basic sentence segmentation\n    sentences = re.split(\n        r\"(?<=[.!?])\\s+|;\\s+|\\|\\|\",\n        text\n    )\n\n    sentences = [\n        s.strip()\n        for s in sentences\n        if s.strip()\n    ]\n\n    return sentences\n\n\n# ----------------------------------------------------------------\n# 6. Evidence classifier\n# ----------------------------------------------------------------\n\ndef classify_sentence(sentence, target):\n    sentence_lower = sentence.lower()\n\n    aliases = TARGET_ALIASES[target]\n\n    target_found = any(\n        re.search(pattern, sentence_lower, flags=re.IGNORECASE)\n        for pattern in aliases\n    )\n\n    if not target_found:\n        return \"NO_TARGET_MENTION\"\n\n    # Check nearby negative context\n    negative_found = any(\n        re.search(pattern, sentence_lower, flags=re.IGNORECASE)\n        for pattern in NEGATIVE_PATTERNS\n    )\n\n    # Check positive/pathological context\n    positive_found = any(\n        re.search(pattern, sentence_lower, flags=re.IGNORECASE)\n        for pattern in POSITIVE_PATTERNS\n    )\n\n    # Strong negative constructions take precedence\n    if negative_found:\n        return \"NEGATIVE\"\n\n    if positive_found:\n        return \"POSITIVE\"\n\n    return \"UNCERTAIN\"\n\n\n# ----------------------------------------------------------------\n# 7. Audit explicitly labeled studies\n# ----------------------------------------------------------------\n\nlabel_mask = train[TARGETS].notna().all(axis=1)\nlabeled_train = train.loc[label_mask].copy()\n\nrecords = []\n\nfor _, row in labeled_train.iterrows():\n\n    study_id = row[\"StudyInstanceUID\"]\n    report = row[\"Report\"]\n\n    sentences = split_sentences(report)\n\n    for target in TARGETS:\n\n        label = int(row[target])\n\n        target_results = []\n\n        for sentence in sentences:\n            result = classify_sentence(sentence, target)\n\n            if result != \"NO_TARGET_MENTION\":\n                target_results.append(\n                    {\n                        \"sentence\": sentence,\n                        \"evidence\": result\n                    }\n                )\n\n        # --------------------------------------------------------\n        # Aggregate sentence-level evidence\n        # --------------------------------------------------------\n\n        if not target_results:\n            overall_evidence = \"NO_EVIDENCE\"\n\n        else:\n            evidence_types = [\n                x[\"evidence\"]\n                for x in target_results\n            ]\n\n            if \"POSITIVE\" in evidence_types and \"NEGATIVE\" in evidence_types:\n                overall_evidence = \"CONFLICTING_EVIDENCE\"\n\n            elif \"POSITIVE\" in evidence_types:\n                overall_evidence = \"POSITIVE\"\n\n            elif \"NEGATIVE\" in evidence_types:\n                overall_evidence = \"NEGATIVE\"\n\n            else:\n                overall_evidence = \"UNCERTAIN\"\n\n        # --------------------------------------------------------\n        # Compare against actual dataset label\n        # --------------------------------------------------------\n\n        if overall_evidence == \"POSITIVE\":\n            if label == 1:\n                agreement = \"AGREEMENT_POSITIVE\"\n            else:\n                agreement = \"LABEL_0_REPORT_POSITIVE\"\n\n        elif overall_evidence == \"NEGATIVE\":\n            if label == 0:\n                agreement = \"AGREEMENT_NEGATIVE\"\n            else:\n                agreement = \"LABEL_1_REPORT_NEGATIVE\"\n\n        elif overall_evidence == \"CONFLICTING_EVIDENCE\":\n            agreement = \"CONFLICTING_REPORT\"\n\n        else:\n            agreement = \"UNCLEAR\"\n\n        records.append(\n            {\n                \"StudyInstanceUID\": study_id,\n                \"target\": target,\n                \"label\": label,\n                \"report_evidence\": overall_evidence,\n                \"agreement\": agreement,\n                \"evidence_sentence_count\": len(target_results),\n                \"evidence_text\": \" || \".join(\n                    x[\"sentence\"]\n                    for x in target_results\n                )\n            }\n        )\n\n\nevidence_df = pd.DataFrame(records)\n\n\n# ----------------------------------------------------------------\n# 8. Print overall summary\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"OVERALL EVIDENCE SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    evidence_df[\"report_evidence\"]\n    .value_counts()\n    .to_string()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABEL / REPORT AGREEMENT\")\nprint(\"=\" * 70)\n\nprint(\n    evidence_df[\"agreement\"]\n    .value_counts()\n    .to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 9. Per-target summary\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PER-TARGET SUMMARY\")\nprint(\"=\" * 70)\n\ntarget_summary = (\n    evidence_df\n    .groupby([\"target\", \"agreement\"])\n    .size()\n    .unstack(fill_value=0)\n)\n\nprint(target_summary)\n\n\n# ----------------------------------------------------------------\n# 10. Show high-confidence examples\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"SAMPLE POSITIVE AGREEMENTS\")\nprint(\"=\" * 70)\n\npositive_examples = evidence_df[\n    evidence_df[\"agreement\"] == \"AGREEMENT_POSITIVE\"\n].head(20)\n\nif len(positive_examples) > 0:\n    print(\n        positive_examples[\n            [\n                \"target\",\n                \"label\",\n                \"evidence_text\"\n            ]\n        ].to_string(index=False)\n    )\nelse:\n    print(\"No positive agreements found.\")\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"SAMPLE NEGATIVE AGREEMENTS\")\nprint(\"=\" * 70)\n\nnegative_examples = evidence_df[\n    evidence_df[\"agreement\"] == \"AGREEMENT_NEGATIVE\"\n].head(20)\n\nif len(negative_examples) > 0:\n    print(\n        negative_examples[\n            [\n                \"target\",\n                \"label\",\n                \"evidence_text\"\n            ]\n        ].to_string(index=False)\n    )\nelse:\n    print(\"No negative agreements found.\")\n\n\n# ----------------------------------------------------------------\n# 11. Show conflicts\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABEL 0 / REPORT POSITIVE\")\nprint(\"=\" * 70)\n\nconflict_0 = evidence_df[\n    evidence_df[\"agreement\"] == \"LABEL_0_REPORT_POSITIVE\"\n]\n\nprint(f\"Count: {len(conflict_0)}\")\n\nif len(conflict_0) > 0:\n    print(\n        conflict_0[\n            [\n                \"StudyInstanceUID\",\n                \"target\",\n                \"label\",\n                \"report_evidence\",\n                \"evidence_text\"\n            ]\n        ].head(30).to_string(index=False)\n    )\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABEL 1 / REPORT NEGATIVE\")\nprint(\"=\" * 70)\n\nconflict_1 = evidence_df[\n    evidence_df[\"agreement\"] == \"LABEL_1_REPORT_NEGATIVE\"\n]\n\nprint(f\"Count: {len(conflict_1)}\")\n\nif len(conflict_1) > 0:\n    print(\n        conflict_1[\n            [\n                \"StudyInstanceUID\",\n                \"target\",\n                \"label\",\n                \"report_evidence\",\n                \"evidence_text\"\n            ]\n        ].to_string(index=False)\n    )\n\n\n# ----------------------------------------------------------------\n# 12. Save outputs\n# ----------------------------------------------------------------\n\nOUTPUT_DIR = \"/kaggle/working/rsna_knee_audit\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\nevidence_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell23_improved_report_evidence.csv\"\n)\n\nsummary_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell23_improved_report_evidence_summary.csv\"\n)\n\nevidence_df.to_csv(\n    evidence_path,\n    index=False\n)\n\ntarget_summary.to_csv(\n    summary_path\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 23 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(f\"Saved:\")\nprint(evidence_path)\nprint(summary_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:38:37.168729Z","iopub.execute_input":"2026-08-11T08:38:37.169121Z","iopub.status.idle":"2026-08-11T08:38:37.342981Z","shell.execute_reply.started":"2026-08-11T08:38:37.169091Z","shell.execute_reply":"2026-08-11T08:38:37.341684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 24 - TRAIN / TEST LEAKAGE AND STUDY INTEGRITY AUDIT\n# ================================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 70)\nprint(\"TRAIN / TEST LEAKAGE AND STUDY INTEGRITY AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# 1. Load required CSVs if necessary\n# ----------------------------------------------------------------\n\nBASE_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\ntrain_csv = os.path.join(BASE_DIR, \"train.csv\")\ntest_csv = os.path.join(BASE_DIR, \"test.csv\")\ntrain_series_csv = os.path.join(BASE_DIR, \"train_series.csv\")\ntest_series_csv = os.path.join(BASE_DIR, \"test_series.csv\")\n\nif \"train\" not in globals():\n    train = pd.read_csv(train_csv)\n\nif \"test\" not in globals():\n    test = pd.read_csv(test_csv)\n\nif \"train_series\" not in globals():\n    train_series = pd.read_csv(train_series_csv)\n\nif \"test_series\" not in globals():\n    test_series = pd.read_csv(test_series_csv)\n\nprint(f\"Train studies: {len(train)}\")\nprint(f\"Test studies: {len(test)}\")\nprint(f\"Train series rows: {len(train_series)}\")\nprint(f\"Test series rows: {len(test_series)}\")\n\n\n# ----------------------------------------------------------------\n# 2. StudyInstanceUID overlap\n# ----------------------------------------------------------------\n\ntrain_studies = set(\n    train[\"StudyInstanceUID\"].astype(str)\n)\n\ntest_studies = set(\n    test[\"StudyInstanceUID\"].astype(str)\n)\n\nstudy_overlap = train_studies.intersection(test_studies)\n\nprint()\nprint(\"=\" * 70)\nprint(\"STUDY INSTANCE UID OVERLAP\")\nprint(\"=\" * 70)\n\nprint(f\"Train unique studies: {len(train_studies)}\")\nprint(f\"Test unique studies: {len(test_studies)}\")\nprint(f\"Overlapping studies: {len(study_overlap)}\")\n\nif study_overlap:\n    print(\"WARNING: StudyInstanceUID leakage detected.\")\n    print(list(study_overlap)[:20])\nelse:\n    print(\"PASS: No StudyInstanceUID overlap.\")\n\n\n# ----------------------------------------------------------------\n# 3. SeriesInstanceUID overlap\n# ----------------------------------------------------------------\n\ntrain_series_ids = set(\n    train_series[\"SeriesInstanceUID\"].astype(str)\n)\n\ntest_series_ids = set(\n    test_series[\"SeriesInstanceUID\"].astype(str)\n)\n\nseries_overlap = train_series_ids.intersection(\n    test_series_ids\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"SERIES INSTANCE UID OVERLAP\")\nprint(\"=\" * 70)\n\nprint(f\"Train unique series: {len(train_series_ids)}\")\nprint(f\"Test unique series: {len(test_series_ids)}\")\nprint(f\"Overlapping series: {len(series_overlap)}\")\n\nif series_overlap:\n    print(\"WARNING: SeriesInstanceUID leakage detected.\")\n    print(list(series_overlap)[:20])\nelse:\n    print(\"PASS: No SeriesInstanceUID overlap.\")\n\n\n# ----------------------------------------------------------------\n# 4. Verify every train series belongs to train\n# ----------------------------------------------------------------\n\ntrain_study_ids_from_series = set(\n    train_series[\"StudyInstanceUID\"].astype(str)\n)\n\ntest_study_ids_from_series = set(\n    test_series[\"StudyInstanceUID\"].astype(str)\n)\n\ntrain_series_orphan_studies = (\n    train_study_ids_from_series - train_studies\n)\n\ntest_series_orphan_studies = (\n    test_study_ids_from_series - test_studies\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"SERIES -> STUDY CONSISTENCY\")\nprint(\"=\" * 70)\n\nprint(\n    \"Train-series studies not present in train.csv:\",\n    len(train_series_orphan_studies)\n)\n\nprint(\n    \"Test-series studies not present in test.csv:\",\n    len(test_series_orphan_studies)\n)\n\nif len(train_series_orphan_studies) == 0:\n    print(\"PASS: Every train series maps to a train study.\")\n\nif len(test_series_orphan_studies) == 0:\n    print(\"PASS: Every test series maps to a test study.\")\n\n\n# ----------------------------------------------------------------\n# 5. Check series counts per study\n# ----------------------------------------------------------------\n\ntrain_series_counts = (\n    train_series\n    .groupby(\"StudyInstanceUID\")[\"SeriesInstanceUID\"]\n    .nunique()\n)\n\ntest_series_counts = (\n    test_series\n    .groupby(\"StudyInstanceUID\")[\"SeriesInstanceUID\"]\n    .nunique()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"SERIES PER STUDY\")\nprint(\"=\" * 70)\n\nprint(\"TRAIN:\")\nprint(train_series_counts.describe().to_string())\n\nprint()\nprint(\"TEST:\")\nprint(test_series_counts.describe().to_string())\n\n\n# ----------------------------------------------------------------\n# 6. Check duplicate StudyInstanceUID rows\n# ----------------------------------------------------------------\n\ntrain_duplicate_rows = train[\n    train[\"StudyInstanceUID\"].duplicated(keep=False)\n]\n\ntest_duplicate_rows = test[\n    test[\"StudyInstanceUID\"].duplicated(keep=False)\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"DUPLICATE STUDY ROWS\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Duplicate train study rows: {len(train_duplicate_rows)}\"\n)\n\nprint(\n    f\"Duplicate test study rows: {len(test_duplicate_rows)}\"\n)\n\n\n# ----------------------------------------------------------------\n# 7. Check duplicate SeriesInstanceUID rows\n# ----------------------------------------------------------------\n\ntrain_series_duplicate_rows = train_series[\n    train_series[\"SeriesInstanceUID\"].duplicated(keep=False)\n]\n\ntest_series_duplicate_rows = test_series[\n    test_series[\"SeriesInstanceUID\"].duplicated(keep=False)\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"DUPLICATE SERIES ROWS\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Duplicate train series rows: {len(train_series_duplicate_rows)}\"\n)\n\nprint(\n    f\"Duplicate test series rows: {len(test_series_duplicate_rows)}\"\n)\n\n\n# ----------------------------------------------------------------\n# 8. Verify StudyInstanceUID -> series mapping\n# ----------------------------------------------------------------\n\ntrain_mapping_counts = (\n    train_series\n    .groupby(\"SeriesInstanceUID\")[\"StudyInstanceUID\"]\n    .nunique()\n)\n\ntest_mapping_counts = (\n    test_series\n    .groupby(\"SeriesInstanceUID\")[\"StudyInstanceUID\"]\n    .nunique()\n)\n\ntrain_bad_series_mapping = (\n    train_mapping_counts[\n        train_mapping_counts != 1\n    ]\n)\n\ntest_bad_series_mapping = (\n    test_mapping_counts[\n        test_mapping_counts != 1\n    ]\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"SERIES -> STUDY ONE-TO-ONE MAPPING\")\nprint(\"=\" * 70)\n\nprint(\n    \"Train series with multiple studies:\",\n    len(train_bad_series_mapping)\n)\n\nprint(\n    \"Test series with multiple studies:\",\n    len(test_bad_series_mapping)\n)\n\nif len(train_bad_series_mapping) == 0:\n    print(\"PASS: Every train series maps to exactly one study.\")\n\nif len(test_bad_series_mapping) == 0:\n    print(\"PASS: Every test series maps to exactly one study.\")\n\n\n# ----------------------------------------------------------------\n# 9. Labelled studies and series\n# ----------------------------------------------------------------\n\nif \"TARGETS\" not in globals():\n    TARGETS = [\n        \"ACL\",\n        \"MCL\",\n        \"Medial Meniscus\",\n        \"Lateral Meniscus\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Synovitis\",\n        \"Baker's\",\n        \"Contusion\",\n        \"Fracture\"\n    ]\n\nlabeled_mask = train[TARGETS].notna().all(axis=1)\n\nlabeled_studies = set(\n    train.loc[\n        labeled_mask,\n        \"StudyInstanceUID\"\n    ].astype(str)\n)\n\nlabeled_series = train_series[\n    train_series[\"StudyInstanceUID\"].astype(str).isin(\n        labeled_studies\n    )\n].copy()\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABELED STUDY COVERAGE\")\nprint(\"=\" * 70)\n\nprint(f\"Fully labeled studies: {len(labeled_studies)}\")\nprint(f\"Series belonging to labeled studies: {len(labeled_series)}\")\nprint(\n    f\"Unique series in labeled studies: \"\n    f\"{labeled_series['SeriesInstanceUID'].nunique()}\"\n)\n\n\n# ----------------------------------------------------------------\n# 10. Metadata overlap between train and test\n# ----------------------------------------------------------------\n\nmetadata_columns = [\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\",\n    \"Anatomical_Plane\"\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"TRAIN / TEST SERIES METADATA DISTRIBUTION\")\nprint(\"=\" * 70)\n\nfor col in metadata_columns:\n\n    print()\n    print(f\"--- {col} ---\")\n\n    train_dist = (\n        train_series[col]\n        .value_counts(dropna=False)\n        .rename(\"train_count\")\n    )\n\n    test_dist = (\n        test_series[col]\n        .value_counts(dropna=False)\n        .rename(\"test_count\")\n    )\n\n    comparison = pd.concat(\n        [\n            train_dist,\n            test_dist\n        ],\n        axis=1\n    ).fillna(0)\n\n    print(comparison)\n\n\n# ----------------------------------------------------------------\n# 11. Final integrity verdict\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL INTEGRITY VERDICT\")\nprint(\"=\" * 70)\n\nchecks = {\n    \"Study UID leakage\": len(study_overlap) == 0,\n    \"Series UID leakage\": len(series_overlap) == 0,\n    \"Train series orphan studies\": len(train_series_orphan_studies) == 0,\n    \"Test series orphan studies\": len(test_series_orphan_studies) == 0,\n    \"Duplicate train studies\": len(train_duplicate_rows) == 0,\n    \"Duplicate test studies\": len(test_duplicate_rows) == 0,\n    \"Duplicate train series\": len(train_series_duplicate_rows) == 0,\n    \"Duplicate test series\": len(test_series_duplicate_rows) == 0,\n    \"Train series mapping\": len(train_bad_series_mapping) == 0,\n    \"Test series mapping\": len(test_bad_series_mapping) == 0\n}\n\nfor check_name, passed in checks.items():\n    print(\n        f\"{'PASS' if passed else 'FAIL'} - {check_name}\"\n    )\n\nall_pass = all(checks.values())\n\nprint()\nprint(\n    \"OVERALL:\",\n    \"PASS\" if all_pass else \"REVIEW REQUIRED\"\n)\n\n\n# ----------------------------------------------------------------\n# 12. Save audit\n# ----------------------------------------------------------------\n\nOUTPUT_DIR = \"/kaggle/working/rsna_knee_audit\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\naudit_rows = []\n\nfor check_name, passed in checks.items():\n    audit_rows.append(\n        {\n            \"check\": check_name,\n            \"passed\": passed\n        }\n    )\n\naudit_df = pd.DataFrame(audit_rows)\n\naudit_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell24_train_test_integrity_audit.csv\"\n)\n\naudit_df.to_csv(\n    audit_path,\n    index=False\n)\n\nprint()\nprint(f\"Saved: {audit_path}\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 24 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:39:39.259355Z","iopub.execute_input":"2026-08-11T08:39:39.259651Z","iopub.status.idle":"2026-08-11T08:39:39.382223Z","shell.execute_reply.started":"2026-08-11T08:39:39.259627Z","shell.execute_reply":"2026-08-11T08:39:39.381352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 25 - LABEL-AWARE VALIDATION FEASIBILITY AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\n\nprint(\"=\" * 70)\nprint(\"LABEL-AWARE VALIDATION FEASIBILITY AUDIT\")\nprint(\"=\" * 70)\n\n\n# ----------------------------------------------------------------\n# 1. Load train data if necessary\n# ----------------------------------------------------------------\n\nBASE_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\ntrain_csv = os.path.join(BASE_DIR, \"train.csv\")\n\nif \"train\" not in globals():\n    train = pd.read_csv(train_csv)\n\nif \"TARGETS\" not in globals():\n    TARGETS = [\n        \"ACL\",\n        \"MCL\",\n        \"Medial Meniscus\",\n        \"Lateral Meniscus\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Synovitis\",\n        \"Baker's\",\n        \"Contusion\",\n        \"Fracture\"\n    ]\n\n\n# ----------------------------------------------------------------\n# 2. Select fully labeled studies\n# ----------------------------------------------------------------\n\nlabel_mask = train[TARGETS].notna().all(axis=1)\n\nlabeled_train = train.loc[label_mask].copy()\n\nprint(f\"Total train studies: {len(train)}\")\nprint(f\"Fully labeled studies: {len(labeled_train)}\")\n\n\n# ----------------------------------------------------------------\n# 3. Basic label counts\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABEL COUNTS\")\nprint(\"=\" * 70)\n\nlabel_counts = []\n\nfor target in TARGETS:\n\n    positive = int(\n        (labeled_train[target] == 1).sum()\n    )\n\n    negative = int(\n        (labeled_train[target] == 0).sum()\n    )\n\n    label_counts.append(\n        {\n            \"target\": target,\n            \"positive\": positive,\n            \"negative\": negative,\n            \"total\": positive + negative,\n            \"positive_rate\": positive / (positive + negative)\n        }\n    )\n\nlabel_counts_df = pd.DataFrame(label_counts)\n\nprint(\n    label_counts_df.to_string(\n        index=False\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 4. Determine feasible number of folds\n# ----------------------------------------------------------------\n\nmin_positive = label_counts_df[\"positive\"].min()\nmin_negative = label_counts_df[\"negative\"].min()\n\nprint()\nprint(\"=\" * 70)\nprint(\"FOLD FEASIBILITY\")\nprint(\"=\" * 70)\n\nprint(f\"Minimum positive count across targets: {min_positive}\")\nprint(f\"Minimum negative count across targets: {min_negative}\")\n\nprint()\nprint(\n    \"Important: For ordinary stratification of a binary target, \"\n    \"the number of folds cannot exceed the smaller class count.\"\n)\n\n\n# ----------------------------------------------------------------\n# 5. Examine candidate fold counts\n# ----------------------------------------------------------------\n\ncandidate_folds = [2, 3, 4, 5]\n\nfold_feasibility = []\n\nfor n_splits in candidate_folds:\n\n    feasible = (\n        min_positive >= n_splits\n        and min_negative >= n_splits\n    )\n\n    fold_feasibility.append(\n        {\n            \"n_splits\": n_splits,\n            \"feasible_for_all_targets\": feasible\n        }\n    )\n\nfold_feasibility_df = pd.DataFrame(\n    fold_feasibility\n)\n\nprint(\n    fold_feasibility_df.to_string(\n        index=False\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 6. Build multilabel stratification signature\n#\n# We create a compact representation of the 12-label combination.\n# This is NOT our final CV splitter. It is only a feasibility\n# diagnostic.\n# ----------------------------------------------------------------\n\nlabel_signature = (\n    labeled_train[TARGETS]\n    .astype(int)\n    .astype(str)\n    .agg(\"\".join, axis=1)\n)\n\nsignature_counts = (\n    label_signature\n    .value_counts()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"MULTILABEL COMBINATION FREQUENCY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Unique 12-target label combinations: \"\n    f\"{len(signature_counts)}\"\n)\n\nprint()\nprint(\"Most common combinations:\")\n\nprint(\n    signature_counts.head(20).to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 7. Check how many unique combinations occur only once\n# ----------------------------------------------------------------\n\nsingletons = int(\n    (signature_counts == 1).sum()\n)\n\nrare_combinations = int(\n    (signature_counts <= 2).sum()\n)\n\nprint()\nprint(\n    f\"Combinations occurring exactly once: {singletons}\"\n)\n\nprint(\n    f\"Combinations occurring at most twice: {rare_combinations}\"\n)\n\n\n# ----------------------------------------------------------------\n# 8. Simulate multilabel-aware fold assignment\n#\n# We use an iterative greedy assignment based on label frequencies.\n# This is a diagnostic splitter, not the final production splitter.\n# ----------------------------------------------------------------\n\ndef iterative_multilabel_split(\n    labels,\n    n_splits=3,\n    random_state=42\n):\n    \"\"\"\n    Simple iterative multilabel stratification diagnostic.\n\n    Attempts to distribute rare positive labels across folds\n    while keeping fold sizes approximately balanced.\n    \"\"\"\n\n    rng = np.random.default_rng(random_state)\n\n    Y = np.asarray(labels).astype(int)\n\n    n_samples = Y.shape[0]\n\n    folds = [\n        []\n        for _ in range(n_splits)\n    ]\n\n    # Start with rarest positive label\n    remaining = set(range(n_samples))\n\n    label_counts = Y.sum(axis=0)\n\n    # Sort samples by rarity of their positive labels\n    sample_priority = []\n\n    for i in range(n_samples):\n\n        positive_labels = np.where(\n            Y[i] == 1\n        )[0]\n\n        if len(positive_labels) == 0:\n            rarity_score = 0\n        else:\n            rarity_score = sum(\n                1.0 / max(label_counts[j], 1)\n                for j in positive_labels\n            )\n\n        sample_priority.append(\n            rarity_score\n        )\n\n    order = np.argsort(\n        -np.asarray(sample_priority)\n    )\n\n    # Assign samples greedily\n    fold_label_counts = np.zeros(\n        (n_splits, Y.shape[1]),\n        dtype=int\n    )\n\n    fold_sizes = np.zeros(\n        n_splits,\n        dtype=int\n    )\n\n    for idx in order:\n\n        sample = Y[idx]\n\n        scores = []\n\n        for fold_idx in range(n_splits):\n\n            label_score = np.sum(\n                fold_label_counts[fold_idx]\n                * sample\n            )\n\n            size_score = fold_sizes[fold_idx]\n\n            scores.append(\n                label_score * 10.0\n                + size_score\n            )\n\n        min_score = min(scores)\n\n        candidates = [\n            i\n            for i, score in enumerate(scores)\n            if score == min_score\n        ]\n\n        chosen = rng.choice(\n            candidates\n        )\n\n        folds[chosen].append(idx)\n\n        fold_label_counts[\n            chosen\n        ] += sample\n\n        fold_sizes[\n            chosen\n        ] += 1\n\n    return folds\n\n\n# ----------------------------------------------------------------\n# 9. Analyze candidate fold structures\n# ----------------------------------------------------------------\n\nfold_diagnostics = []\n\nfor n_splits in [3, 4, 5]:\n\n    folds = iterative_multilabel_split(\n        labeled_train[TARGETS].values,\n        n_splits=n_splits,\n        random_state=42\n    )\n\n    print()\n    print(\"=\" * 70)\n    print(f\"{n_splits}-FOLD MULTILABEL DIAGNOSTIC\")\n    print(\"=\" * 70)\n\n    for fold_idx, indices in enumerate(folds):\n\n        fold_data = labeled_train.iloc[\n            indices\n        ]\n\n        print()\n        print(\n            f\"Fold {fold_idx + 1}: \"\n            f\"{len(fold_data)} studies\"\n        )\n\n        for target in TARGETS:\n\n            positives = int(\n                (fold_data[target] == 1).sum()\n            )\n\n            negatives = int(\n                (fold_data[target] == 0).sum()\n            )\n\n            print(\n                f\"{target:20s} \"\n                f\"pos={positives:2d} \"\n                f\"neg={negatives:2d}\"\n            )\n\n            fold_diagnostics.append(\n                {\n                    \"n_splits\": n_splits,\n                    \"fold\": fold_idx + 1,\n                    \"target\": target,\n                    \"positive\": positives,\n                    \"negative\": negatives\n                }\n            )\n\n\n# ----------------------------------------------------------------\n# 10. Identify impossible AUC folds\n# ----------------------------------------------------------------\n\ndiagnostics_df = pd.DataFrame(\n    fold_diagnostics\n)\n\ndiagnostics_df[\"auc_possible\"] = (\n    (diagnostics_df[\"positive\"] > 0)\n    &\n    (diagnostics_df[\"negative\"] > 0)\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"AUC FEASIBILITY\")\nprint(\"=\" * 70)\n\nfor n_splits in [3, 4, 5]:\n\n    subset = diagnostics_df[\n        diagnostics_df[\"n_splits\"] == n_splits\n    ]\n\n    impossible = subset[\n        ~subset[\"auc_possible\"]\n    ]\n\n    print(\n        f\"{n_splits}-fold: \"\n        f\"{len(impossible)} target-fold combinations \"\n        f\"without both classes\"\n    )\n\n    if len(impossible) > 0:\n        print(\n            impossible[\n                [\n                    \"fold\",\n                    \"target\",\n                    \"positive\",\n                    \"negative\"\n                ]\n            ].to_string(index=False)\n        )\n\n\n# ----------------------------------------------------------------\n# 11. Estimate minimum positive examples per validation fold\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"MINIMUM POSITIVE COUNT PER FOLD\")\nprint(\"=\" * 70)\n\nfor n_splits in [3, 4, 5]:\n\n    subset = diagnostics_df[\n        diagnostics_df[\"n_splits\"] == n_splits\n    ]\n\n    minimum_positive = (\n        subset[\"positive\"].min()\n    )\n\n    minimum_negative = (\n        subset[\"negative\"].min()\n    )\n\n    print(\n        f\"{n_splits}-fold -> \"\n        f\"minimum positive = {minimum_positive}, \"\n        f\"minimum negative = {minimum_negative}\"\n    )\n\n\n# ----------------------------------------------------------------\n# 12. Save outputs\n# ----------------------------------------------------------------\n\nOUTPUT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nos.makedirs(\n    OUTPUT_DIR,\n    exist_ok=True\n)\n\nlabel_counts_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell25_label_counts.csv\"\n)\n\nfold_feasibility_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell25_fold_feasibility.csv\"\n)\n\nsignature_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell25_label_signature_counts.csv\"\n)\n\ndiagnostics_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell25_fold_diagnostics.csv\"\n)\n\nlabel_counts_df.to_csv(\n    label_counts_path,\n    index=False\n)\n\nfold_feasibility_df.to_csv(\n    fold_feasibility_path,\n    index=False\n)\n\nsignature_counts.rename(\n    \"count\"\n).to_csv(\n    signature_path\n)\n\ndiagnostics_df.to_csv(\n    diagnostics_path,\n    index=False\n)\n\n\n# ----------------------------------------------------------------\n# 13. Final message\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 25 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"Saved:\")\nprint(label_counts_path)\nprint(fold_feasibility_path)\nprint(signature_path)\nprint(diagnostics_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:40:56.94178Z","iopub.execute_input":"2026-08-11T08:40:56.942152Z","iopub.status.idle":"2026-08-11T08:40:58.216362Z","shell.execute_reply.started":"2026-08-11T08:40:56.942123Z","shell.execute_reply":"2026-08-11T08:40:58.215276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 26 - LABELED STUDY MRI SEQUENCE / SERIES AUDIT\n# ================================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 70)\nprint(\"LABELED STUDY MRI SEQUENCE / SERIES AUDIT\")\nprint(\"=\" * 70)\n\n\n# ----------------------------------------------------------------\n# 1. Load required data\n# ----------------------------------------------------------------\n\nBASE_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nif \"train\" not in globals():\n    train = pd.read_csv(\n        os.path.join(BASE_DIR, \"train.csv\")\n    )\n\nif \"train_series\" not in globals():\n    train_series = pd.read_csv(\n        os.path.join(BASE_DIR, \"train_series.csv\")\n    )\n\nif \"TARGETS\" not in globals():\n    TARGETS = [\n        \"ACL\",\n        \"MCL\",\n        \"Medial Meniscus\",\n        \"Lateral Meniscus\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Synovitis\",\n        \"Baker's\",\n        \"Contusion\",\n        \"Fracture\"\n    ]\n\n\n# ----------------------------------------------------------------\n# 2. Select the 58 fully labeled studies\n# ----------------------------------------------------------------\n\nlabel_mask = train[TARGETS].notna().all(axis=1)\n\nlabeled_train = train.loc[\n    label_mask\n].copy()\n\nlabeled_study_ids = set(\n    labeled_train[\n        \"StudyInstanceUID\"\n    ].astype(str)\n)\n\nlabeled_series = train_series[\n    train_series[\n        \"StudyInstanceUID\"\n    ].astype(str).isin(\n        labeled_study_ids\n    )\n].copy()\n\nprint(\n    f\"Fully labeled studies: \"\n    f\"{len(labeled_train)}\"\n)\n\nprint(\n    f\"Labeled-study series rows: \"\n    f\"{len(labeled_series)}\"\n)\n\nprint(\n    f\"Unique labeled-study series: \"\n    f\"{labeled_series['SeriesInstanceUID'].nunique()}\"\n)\n\n\n# ----------------------------------------------------------------\n# 3. Series per labeled study\n# ----------------------------------------------------------------\n\nseries_per_study = (\n    labeled_series\n    .groupby(\"StudyInstanceUID\")\n    [\"SeriesInstanceUID\"]\n    .nunique()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"SERIES PER LABELED STUDY\")\nprint(\"=\" * 70)\n\nprint(\n    series_per_study.describe().to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 4. Plane distribution\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABELED-STUDY PLANE DISTRIBUTION\")\nprint(\"=\" * 70)\n\nplane_counts = (\n    labeled_series[\n        \"Anatomical_Plane\"\n    ]\n    .value_counts(dropna=False)\n)\n\nplane_percent = (\n    plane_counts\n    / plane_counts.sum()\n    * 100\n)\n\nplane_table = pd.DataFrame({\n    \"count\": plane_counts,\n    \"percentage\": plane_percent.round(2)\n})\n\nprint(\n    plane_table.to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 5. Sequence property distribution\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FLUID-SENSITIVE DISTRIBUTION\")\nprint(\"=\" * 70)\n\nfluid_counts = (\n    labeled_series[\n        \"Fluid_Sensitive\"\n    ]\n    .value_counts(dropna=False)\n)\n\nfluid_table = pd.DataFrame({\n    \"count\": fluid_counts,\n    \"percentage\": (\n        fluid_counts\n        / fluid_counts.sum()\n        * 100\n    ).round(2)\n})\n\nprint(\n    fluid_table.to_string()\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"FAT-SUPPRESSION DISTRIBUTION\")\nprint(\"=\" * 70)\n\nfat_counts = (\n    labeled_series[\n        \"Fat_Suppression\"\n    ]\n    .value_counts(dropna=False)\n)\n\nfat_table = pd.DataFrame({\n    \"count\": fat_counts,\n    \"percentage\": (\n        fat_counts\n        / fat_counts.sum()\n        * 100\n    ).round(2)\n})\n\nprint(\n    fat_table.to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 6. Plane + sequence combinations\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PLANE + SEQUENCE COMBINATIONS\")\nprint(\"=\" * 70)\n\ncombo = (\n    labeled_series\n    .groupby(\n        [\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\"\n        ],\n        dropna=False\n    )\n    .size()\n    .reset_index(\n        name=\"series_count\"\n    )\n    .sort_values(\n        \"series_count\",\n        ascending=False\n    )\n)\n\ncombo[\"percentage\"] = (\n    combo[\"series_count\"]\n    / len(labeled_series)\n    * 100\n).round(2)\n\nprint(\n    combo.to_string(index=False)\n)\n\n\n# ----------------------------------------------------------------\n# 7. Series descriptions\n#\n# train_series.csv does not necessarily contain descriptions.\n# Therefore we only inspect columns that actually exist.\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"AVAILABLE TRAIN SERIES COLUMNS\")\nprint(\"=\" * 70)\n\nprint(\n    list(train_series.columns)\n)\n\nif \"SeriesDescription\" in labeled_series.columns:\n\n    print()\n    print(\"=\" * 70)\n    print(\"TOP SERIES DESCRIPTIONS\")\n    print(\"=\" * 70)\n\n    descriptions = (\n        labeled_series[\n            \"SeriesDescription\"\n        ]\n        .fillna(\"MISSING\")\n        .astype(str)\n        .value_counts()\n        .head(50)\n    )\n\n    print(\n        descriptions.to_string()\n    )\n\nelse:\n\n    print()\n    print(\n        \"SeriesDescription is not available \"\n        \"in train_series.csv.\"\n    )\n\n\n# ----------------------------------------------------------------\n# 8. Per-study plane coverage\n# ----------------------------------------------------------------\n\nstudy_plane = (\n    labeled_series\n    .pivot_table(\n        index=\"StudyInstanceUID\",\n        columns=\"Anatomical_Plane\",\n        values=\"SeriesInstanceUID\",\n        aggfunc=\"nunique\",\n        fill_value=0\n    )\n)\n\nfor plane in [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]:\n    if plane not in study_plane.columns:\n        study_plane[plane] = 0\n\nstudy_plane = study_plane[\n    [\n        \"Sagittal\",\n        \"Coronal\",\n        \"Axial\"\n    ]\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"PLANE COVERAGE PER LABELED STUDY\")\nprint(\"=\" * 70)\n\nprint(\n    study_plane.describe().to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 9. Count studies containing each plane\n# ----------------------------------------------------------------\n\nplane_presence = pd.DataFrame(\n    {\n        \"plane\": [\n            \"Sagittal\",\n            \"Coronal\",\n            \"Axial\"\n        ],\n        \"studies_with_plane\": [\n            int(\n                (study_plane[\"Sagittal\"] > 0).sum()\n            ),\n            int(\n                (study_plane[\"Coronal\"] > 0).sum()\n            ),\n            int(\n                (study_plane[\"Axial\"] > 0).sum()\n            )\n        ]\n    }\n)\n\nplane_presence[\"percentage_of_labeled_studies\"] = (\n    plane_presence[\"studies_with_plane\"]\n    / len(labeled_train)\n    * 100\n).round(2)\n\nprint()\nprint(\n    plane_presence.to_string(\n        index=False\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 10. Count studies containing each sequence property\n# ----------------------------------------------------------------\n\nfluid_presence = (\n    labeled_series\n    .groupby(\"StudyInstanceUID\")\n    [\"Fluid_Sensitive\"]\n    .apply(\n        lambda x: int((x == 1).any())\n    )\n)\n\nfat_presence = (\n    labeled_series\n    .groupby(\"StudyInstanceUID\")\n    [\"Fat_Suppression\"]\n    .apply(\n        lambda x: int((x == 1).any())\n    )\n)\n\nsequence_presence = pd.DataFrame({\n    \"Fluid_Sensitive_present\": fluid_presence,\n    \"Fat_Suppression_present\": fat_presence\n})\n\nprint()\nprint(\"=\" * 70)\nprint(\"SEQUENCE PROPERTY PRESENCE PER LABELED STUDY\")\nprint(\"=\" * 70)\n\nprint(\n    sequence_presence.sum().to_string()\n)\n\nprint()\n\nprint(\n    (\n        sequence_presence.mean() * 100\n    ).round(2).to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 11. Studies with complete three-plane coverage\n# ----------------------------------------------------------------\n\ncomplete_three_plane = (\n    (study_plane[\"Sagittal\"] > 0)\n    &\n    (study_plane[\"Coronal\"] > 0)\n    &\n    (study_plane[\"Axial\"] > 0)\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"THREE-PLANE COVERAGE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Studies containing all three planes:\",\n    int(complete_three_plane.sum())\n)\n\nprint(\n    \"Percentage:\",\n    round(\n        complete_three_plane.mean() * 100,\n        2\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 12. Studies with multiple series of same plane\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"MULTIPLE SERIES PER PLANE\")\nprint(\"=\" * 70)\n\nfor plane in [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]:\n\n    values = study_plane[plane]\n\n    print()\n    print(\n        f\"{plane}:\"\n    )\n\n    print(\n        f\"  Studies with 0 series: \"\n        f\"{int((values == 0).sum())}\"\n    )\n\n    print(\n        f\"  Studies with 1 series: \"\n        f\"{int((values == 1).sum())}\"\n    )\n\n    print(\n        f\"  Studies with >1 series: \"\n        f\"{int((values > 1).sum())}\"\n    )\n\n    print(\n        f\"  Maximum series in one study: \"\n        f\"{int(values.max())}\"\n    )\n\n\n# ----------------------------------------------------------------\n# 13. Create study-level sequence summary\n# ----------------------------------------------------------------\n\nstudy_sequence_summary = (\n    labeled_series\n    .groupby(\"StudyInstanceUID\")\n    .agg(\n        total_series=(\n            \"SeriesInstanceUID\",\n            \"nunique\"\n        ),\n        fluid_sensitive_series=(\n            \"Fluid_Sensitive\",\n            lambda x: int((x == 1).sum())\n        ),\n        fat_suppressed_series=(\n            \"Fat_Suppression\",\n            lambda x: int((x == 1).sum())\n        ),\n        sagittal_series=(\n            \"Anatomical_Plane\",\n            lambda x: int((x == \"Sagittal\").sum())\n        ),\n        coronal_series=(\n            \"Anatomical_Plane\",\n            lambda x: int((x == \"Coronal\").sum())\n        ),\n        axial_series=(\n            \"Anatomical_Plane\",\n            lambda x: int((x == \"Axial\").sum())\n        )\n    )\n    .reset_index()\n)\n\n\n# ----------------------------------------------------------------\n# 14. Merge labels with study-level sequence information\n# ----------------------------------------------------------------\n\nstudy_sequence_summary = study_sequence_summary.merge(\n    labeled_train[\n        [\"StudyInstanceUID\"] + TARGETS\n    ],\n    on=\"StudyInstanceUID\",\n    how=\"left\"\n)\n\n\n# ----------------------------------------------------------------\n# 15. Save outputs\n# ----------------------------------------------------------------\n\nOUTPUT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nos.makedirs(\n    OUTPUT_DIR,\n    exist_ok=True\n)\n\nstudy_summary_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell26_labeled_study_sequence_summary.csv\"\n)\n\ncombo_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell26_sequence_combinations.csv\"\n)\n\nplane_presence_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell26_plane_presence.csv\"\n)\n\nstudy_sequence_summary.to_csv(\n    study_summary_path,\n    index=False\n)\n\ncombo.to_csv(\n    combo_path,\n    index=False\n)\n\nplane_presence.to_csv(\n    plane_presence_path,\n    index=False\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 26 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"Saved:\")\nprint(study_summary_path)\nprint(combo_path)\nprint(plane_presence_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:42:31.650844Z","iopub.execute_input":"2026-08-11T08:42:31.651233Z","iopub.status.idle":"2026-08-11T08:42:31.778553Z","shell.execute_reply.started":"2026-08-11T08:42:31.651203Z","shell.execute_reply":"2026-08-11T08:42:31.777579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 27 - MODEL INPUT SERIES SELECTION AUDIT\n# ================================================================\n\nimport os\nimport glob\nimport pandas as pd\nimport numpy as np\nimport pydicom\n\nprint(\"=\" * 70)\nprint(\"MODEL INPUT SERIES SELECTION AUDIT\")\nprint(\"=\" * 70)\n\n\n# ----------------------------------------------------------------\n# 1. Load required data\n# ----------------------------------------------------------------\n\nBASE_DIR = \"/kaggle/input/competitions/rsna-knee-abnormality-detection\"\n\nTRAIN_SERIES_DIR = os.path.join(\n    BASE_DIR,\n    \"train_series\"\n)\n\nTEST_SERIES_DIR = os.path.join(\n    BASE_DIR,\n    \"test_series\"\n)\n\nif \"train\" not in globals():\n    train = pd.read_csv(\n        os.path.join(BASE_DIR, \"train.csv\")\n    )\n\nif \"train_series\" not in globals():\n    train_series = pd.read_csv(\n        os.path.join(BASE_DIR, \"train_series.csv\")\n    )\n\nif \"test\" not in globals():\n    test = pd.read_csv(\n        os.path.join(BASE_DIR, \"test.csv\")\n    )\n\nif \"test_series\" not in globals():\n    test_series = pd.read_csv(\n        os.path.join(BASE_DIR, \"test_series.csv\")\n    )\n\nif \"TARGETS\" not in globals():\n    TARGETS = [\n        \"ACL\",\n        \"MCL\",\n        \"Medial Meniscus\",\n        \"Lateral Meniscus\",\n        \"Medial OA\",\n        \"Lateral OA\",\n        \"PF OA\",\n        \"Effusion\",\n        \"Synovitis\",\n        \"Baker's\",\n        \"Contusion\",\n        \"Fracture\"\n    ]\n\n\n# ----------------------------------------------------------------\n# 2. Identify fully labeled studies\n# ----------------------------------------------------------------\n\nlabel_mask = train[TARGETS].notna().all(axis=1)\n\nlabeled_train = train.loc[\n    label_mask\n].copy()\n\nlabeled_study_ids = set(\n    labeled_train[\n        \"StudyInstanceUID\"\n    ].astype(str)\n)\n\nlabeled_series = train_series[\n    train_series[\n        \"StudyInstanceUID\"\n    ].astype(str).isin(\n        labeled_study_ids\n    )\n].copy()\n\nprint(\n    \"Fully labeled studies:\",\n    len(labeled_study_ids)\n)\n\nprint(\n    \"Labeled series:\",\n    len(labeled_series)\n)\n\n\n# ----------------------------------------------------------------\n# 3. Count DICOM files for each series\n# ----------------------------------------------------------------\n\ndef get_series_dir(\n    root_dir,\n    study_uid,\n    series_uid\n):\n    return os.path.join(\n        root_dir,\n        str(study_uid),\n        str(series_uid)\n    )\n\n\ndef count_dicom_files(\n    series_dir\n):\n    if not os.path.isdir(series_dir):\n        return 0\n\n    return len(\n        glob.glob(\n            os.path.join(\n                series_dir,\n                \"*.dcm\"\n            )\n        )\n    )\n\n\n# ----------------------------------------------------------------\n# 4. Build series-level audit table\n# ----------------------------------------------------------------\n\nrows = []\n\nprint()\nprint(\"Auditing labeled-study series...\")\n\nfor idx, row in labeled_series.iterrows():\n\n    study_uid = str(\n        row[\"StudyInstanceUID\"]\n    )\n\n    series_uid = str(\n        row[\"SeriesInstanceUID\"]\n    )\n\n    series_dir = get_series_dir(\n        TRAIN_SERIES_DIR,\n        study_uid,\n        series_uid\n    )\n\n    slice_count = count_dicom_files(\n        series_dir\n    )\n\n    first_dicom = None\n\n    files = glob.glob(\n        os.path.join(\n            series_dir,\n            \"*.dcm\"\n        )\n    )\n\n    if files:\n        try:\n            first_dicom = pydicom.dcmread(\n                files[0],\n                stop_before_pixels=True\n            )\n        except Exception:\n            first_dicom = None\n\n    if first_dicom is not None:\n\n        rows_count = getattr(\n            first_dicom,\n            \"Rows\",\n            np.nan\n        )\n\n        cols_count = getattr(\n            first_dicom,\n            \"Columns\",\n            np.nan\n        )\n\n        pixel_spacing = getattr(\n            first_dicom,\n            \"PixelSpacing\",\n            [np.nan, np.nan]\n        )\n\n        if pixel_spacing is not None:\n            try:\n                pixel_spacing_x = float(\n                    pixel_spacing[0]\n                )\n                pixel_spacing_y = float(\n                    pixel_spacing[1]\n                )\n            except Exception:\n                pixel_spacing_x = np.nan\n                pixel_spacing_y = np.nan\n        else:\n            pixel_spacing_x = np.nan\n            pixel_spacing_y = np.nan\n\n        slice_thickness = getattr(\n            first_dicom,\n            \"SliceThickness\",\n            np.nan\n        )\n\n        spacing_between = getattr(\n            first_dicom,\n            \"SpacingBetweenSlices\",\n            np.nan\n        )\n\n        series_description = getattr(\n            first_dicom,\n            \"SeriesDescription\",\n            \"\"\n        )\n\n        image_type = getattr(\n            first_dicom,\n            \"ImageType\",\n            \"\"\n        )\n\n        if isinstance(\n            image_type,\n            (list, tuple)\n        ):\n            image_type = \"|\".join(\n                map(str, image_type)\n            )\n\n    else:\n\n        rows_count = np.nan\n        cols_count = np.nan\n        pixel_spacing_x = np.nan\n        pixel_spacing_y = np.nan\n        slice_thickness = np.nan\n        spacing_between = np.nan\n        series_description = \"\"\n        image_type = \"\"\n\n    rows.append(\n        {\n            \"StudyInstanceUID\": study_uid,\n            \"SeriesInstanceUID\": series_uid,\n            \"Anatomical_Plane\": row[\n                \"Anatomical_Plane\"\n            ],\n            \"Fluid_Sensitive\": row[\n                \"Fluid_Sensitive\"\n            ],\n            \"Fat_Suppression\": row[\n                \"Fat_Suppression\"\n            ],\n            \"slice_count\": slice_count,\n            \"Rows\": rows_count,\n            \"Columns\": cols_count,\n            \"PixelSpacingX\": pixel_spacing_x,\n            \"PixelSpacingY\": pixel_spacing_y,\n            \"SliceThickness\": slice_thickness,\n            \"SpacingBetweenSlices\": spacing_between,\n            \"SeriesDescription\": str(\n                series_description\n            ),\n            \"ImageType\": str(\n                image_type\n            )\n        }\n    )\n\n    if (\n        len(rows) % 50 == 0\n    ):\n        print(\n            f\"Processed {len(rows)} / \"\n            f\"{len(labeled_series)}\"\n        )\n\n\nseries_audit = pd.DataFrame(\n    rows\n)\n\nprint()\nprint(\n    \"Completed series audit:\",\n    len(series_audit)\n)\n\n\n# ----------------------------------------------------------------\n# 5. Create candidate ranking\n#\n# Priority:\n#   1. Fluid sensitive\n#   2. Fat suppression\n#   3. Reasonable slice count\n#   4. Higher spatial resolution\n#\n# We do NOT automatically claim this is the final model choice.\n# This is an audit to identify strong candidate series.\n# ----------------------------------------------------------------\n\nseries_audit[\"resolution_score\"] = (\n    series_audit[\"Rows\"].fillna(0)\n    *\n    series_audit[\"Columns\"].fillna(0)\n)\n\nseries_audit[\"sequence_score\"] = (\n    series_audit[\"Fluid_Sensitive\"].fillna(0)\n    +\n    series_audit[\"Fat_Suppression\"].fillna(0)\n)\n\nseries_audit[\"candidate_score\"] = (\n    series_audit[\"sequence_score\"] * 100000000\n    +\n    series_audit[\"resolution_score\"]\n)\n\n\n# ----------------------------------------------------------------\n# 6. Rank candidates within each study and plane\n# ----------------------------------------------------------------\n\nseries_audit = (\n    series_audit\n    .sort_values(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n            \"candidate_score\",\n            \"slice_count\"\n        ],\n        ascending=[\n            True,\n            True,\n            False,\n            False\n        ]\n    )\n)\n\nseries_audit[\"rank_within_plane\"] = (\n    series_audit\n    .groupby(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\"\n        ]\n    )\n    .cumcount()\n    + 1\n)\n\n\n# ----------------------------------------------------------------\n# 7. Show top candidates from labeled studies\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TOP SERIES CANDIDATES\")\nprint(\"=\" * 70)\n\ndisplay_columns = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\",\n    \"Anatomical_Plane\",\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\",\n    \"slice_count\",\n    \"Rows\",\n    \"Columns\",\n    \"PixelSpacingX\",\n    \"PixelSpacingY\",\n    \"SliceThickness\",\n    \"SpacingBetweenSlices\",\n    \"SeriesDescription\",\n    \"rank_within_plane\"\n]\n\nprint(\n    series_audit[\n        display_columns\n    ]\n    .head(40)\n    .to_string(index=False)\n)\n\n\n# ----------------------------------------------------------------\n# 8. Count what the top-ranked candidate looks like\n# ----------------------------------------------------------------\n\ntop_candidates = series_audit[\n    series_audit[\n        \"rank_within_plane\"\n    ] == 1\n].copy()\n\nprint()\nprint(\"=\" * 70)\nprint(\"TOP-RANKED SERIES DISTRIBUTION\")\nprint(\"=\" * 70)\n\nprint(\n    top_candidates[\n        [\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\"\n        ]\n    ]\n    .value_counts()\n    .to_string()\n)\n\n\n# ----------------------------------------------------------------\n# 9. Check whether every labeled study has a candidate\n#    in all three planes\n# ----------------------------------------------------------------\n\ncandidate_plane_counts = (\n    top_candidates\n    .groupby(\n        \"StudyInstanceUID\"\n    )[\"Anatomical_Plane\"]\n    .nunique()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"TOP-CANDIDATE THREE-PLANE COVERAGE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Studies:\",\n    len(candidate_plane_counts)\n)\n\nprint(\n    \"Studies with 3 planes:\",\n    int(\n        (\n            candidate_plane_counts == 3\n        ).sum()\n    )\n)\n\nprint(\n    \"Studies missing at least one plane:\",\n    int(\n        (\n            candidate_plane_counts < 3\n        ).sum()\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 10. Test-set series audit\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST SERIES SUMMARY\")\nprint(\"=\" * 70)\n\ntest_summary = (\n    test_series\n    .groupby(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\"\n        ]\n    )\n    .size()\n    .reset_index(\n        name=\"series_count\"\n    )\n)\n\nprint(\n    test_summary.to_string(\n        index=False\n    )\n)\n\n\n# ----------------------------------------------------------------\n# 11. Save\n# ----------------------------------------------------------------\n\nOUTPUT_DIR = (\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nos.makedirs(\n    OUTPUT_DIR,\n    exist_ok=True\n)\n\nseries_audit_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell27_labeled_series_candidate_audit.csv\"\n)\n\ntop_candidates_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell27_top_series_candidates.csv\"\n)\n\ntest_summary_path = os.path.join(\n    OUTPUT_DIR,\n    \"cell27_test_series_summary.csv\"\n)\n\nseries_audit.to_csv(\n    series_audit_path,\n    index=False\n)\n\ntop_candidates.to_csv(\n    top_candidates_path,\n    index=False\n)\n\ntest_summary.to_csv(\n    test_summary_path,\n    index=False\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 27 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"Saved:\")\nprint(series_audit_path)\nprint(top_candidates_path)\nprint(test_summary_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:43:29.151065Z","iopub.execute_input":"2026-08-11T08:43:29.151425Z","iopub.status.idle":"2026-08-11T08:43:32.967956Z","shell.execute_reply.started":"2026-08-11T08:43:29.151361Z","shell.execute_reply":"2026-08-11T08:43:32.966813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 28 - MODELING ENVIRONMENT AND GPU CHECK\n# ================================================================\n\nimport os\nimport sys\nimport platform\nimport subprocess\nimport importlib.util\n\nprint(\"=\" * 70)\nprint(\"MODELING ENVIRONMENT CHECK\")\nprint(\"=\" * 70)\n\n# ---------------------------------------------------------------\n# Python\n# ---------------------------------------------------------------\n\nprint(\"\\nPython:\")\nprint(sys.version)\n\nprint(\"\\nPlatform:\")\nprint(platform.platform())\n\n\n# ---------------------------------------------------------------\n# PyTorch\n# ---------------------------------------------------------------\n\nprint(\"\\nChecking PyTorch...\")\n\nif importlib.util.find_spec(\"torch\") is None:\n    print(\"ERROR: PyTorch is not installed.\")\nelse:\n    import torch\n\n    print(\"PyTorch version:\", torch.__version__)\n    print(\"CUDA available:\", torch.cuda.is_available())\n\n    if torch.cuda.is_available():\n\n        print(\"CUDA version:\", torch.version.cuda)\n        print(\n            \"GPU count:\",\n            torch.cuda.device_count()\n        )\n\n        for i in range(torch.cuda.device_count()):\n\n            props = torch.cuda.get_device_properties(i)\n\n            print()\n            print(f\"GPU {i}:\")\n            print(\"Name:\", props.name)\n            print(\n                \"VRAM:\",\n                round(\n                    props.total_memory / (1024 ** 3),\n                    2\n                ),\n                \"GB\"\n            )\n\n            print(\n                \"Compute capability:\",\n                f\"{props.major}.{props.minor}\"\n            )\n\n    else:\n        print(\n            \"WARNING: CUDA is not available.\"\n        )\n\n\n# ---------------------------------------------------------------\n# Important packages\n# ---------------------------------------------------------------\n\npackages = [\n    \"numpy\",\n    \"pandas\",\n    \"pydicom\",\n    \"sklearn\",\n    \"PIL\",\n    \"torch\",\n    \"torchvision\"\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"PACKAGE CHECK\")\nprint(\"=\" * 70)\n\nfor package in packages:\n\n    installed = (\n        importlib.util.find_spec(package)\n        is not None\n    )\n\n    print(\n        f\"{package:15s}: \"\n        f\"{'AVAILABLE' if installed else 'MISSING'}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Output/storage check\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"STORAGE CHECK\")\nprint(\"=\" * 70)\n\npaths_to_check = [\n    \"/kaggle/input\",\n    \"/kaggle/working\",\n    \"/kaggle/working/rsna_knee_audit\"\n]\n\nfor path in paths_to_check:\n\n    if os.path.exists(path):\n\n        try:\n\n            usage = (\n                subprocess.check_output(\n                    [\"df\", \"-h\", path],\n                    text=True\n                )\n                .strip()\n                .split(\"\\n\")\n            )\n\n            print()\n            print(path)\n            print(usage[-1])\n\n        except Exception as e:\n\n            print(\n                path,\n                \"exists, but disk usage \"\n                \"could not be read.\"\n            )\n\n    else:\n\n        print(\n            path,\n            \"does not exist.\"\n        )\n\n\n# ---------------------------------------------------------------\n# Existing audit files\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"AUDIT FILES\")\nprint(\"=\" * 70)\n\nAUDIT_DIR = (\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nif os.path.isdir(AUDIT_DIR):\n\n    audit_files = sorted(\n        os.listdir(AUDIT_DIR)\n    )\n\n    print(\n        \"Audit files found:\",\n        len(audit_files)\n    )\n\n    for filename in audit_files:\n\n        print(\n            \" -\",\n            filename\n        )\n\nelse:\n\n    print(\n        \"Audit directory not found.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Final verdict\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 28 COMPLETE\")\nprint(\"=\" * 70)\n\nif (\n    importlib.util.find_spec(\"torch\")\n    is not None\n):\n\n    import torch\n\n    if torch.cuda.is_available():\n\n        print(\n            \"READY: CUDA GPU is available.\"\n        )\n\n    else:\n\n        print(\n            \"WARNING: CUDA GPU is not available.\"\n        )\n\nelse:\n\n    print(\n        \"STOP: PyTorch is missing.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:44:53.640288Z","iopub.execute_input":"2026-08-11T08:44:53.640624Z","iopub.status.idle":"2026-08-11T08:44:57.862975Z","shell.execute_reply.started":"2026-08-11T08:44:53.640597Z","shell.execute_reply":"2026-08-11T08:44:57.862053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 29 - BUILD STUDY-LEVEL MODELING TABLE + 3-FOLD SPLIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nfrom collections import defaultdict\n\n# ---------------------------------------------------------------\n# Paths\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nTRAIN_CSV = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/train.csv\"\n)\n\nLABELED_CSV = os.path.join(\n    AUDIT_DIR,\n    \"explicitly_labeled_studies.csv\"\n)\n\nTOP_CANDIDATES_CSV = os.path.join(\n    AUDIT_DIR,\n    \"cell27_top_series_candidates.csv\"\n)\n\nOUTPUT_MODEL_TABLE = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nOUTPUT_FOLDS = os.path.join(\n    AUDIT_DIR,\n    \"cell29_study_folds.csv\"\n)\n\n# ---------------------------------------------------------------\n# Target columns\n# ---------------------------------------------------------------\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\nSEED = 42\nN_FOLDS = 3\n\nprint(\"=\" * 70)\nprint(\"CELL 29 - STUDY-LEVEL MODELING TABLE\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load files\n# ---------------------------------------------------------------\n\nprint(\"\\nLoading training data...\")\n\ntrain = pd.read_csv(TRAIN_CSV)\nlabeled = pd.read_csv(LABELED_CSV)\ncandidates = pd.read_csv(TOP_CANDIDATES_CSV)\n\nprint(\"Train rows:\", len(train))\nprint(\"Labeled rows:\", len(labeled))\nprint(\"Candidate rows:\", len(candidates))\n\n\n# ---------------------------------------------------------------\n# Basic validation\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BASIC INPUT VALIDATION\")\nprint(\"=\" * 70)\n\nmissing_targets = [\n    col for col in TARGETS\n    if col not in train.columns\n]\n\nif missing_targets:\n    raise ValueError(\n        f\"Missing target columns: {missing_targets}\"\n    )\n\nif \"StudyInstanceUID\" not in train.columns:\n    raise ValueError(\n        \"StudyInstanceUID missing from train.csv\"\n    )\n\nif \"StudyInstanceUID\" not in candidates.columns:\n    raise ValueError(\n        \"StudyInstanceUID missing from candidate file\"\n    )\n\nif \"SeriesInstanceUID\" not in candidates.columns:\n    raise ValueError(\n        \"SeriesInstanceUID missing from candidate file\"\n    )\n\nprint(\"All target columns present.\")\nprint(\"StudyInstanceUID present.\")\nprint(\"SeriesInstanceUID present.\")\n\n\n# ---------------------------------------------------------------\n# Restrict to explicitly labeled studies\n# ---------------------------------------------------------------\n\nlabeled_ids = (\n    labeled[\"StudyInstanceUID\"]\n    .astype(str)\n    .unique()\n)\n\nprint()\nprint(\"Explicitly labeled studies:\", len(labeled_ids))\n\nlabeled_train = train[\n    train[\"StudyInstanceUID\"]\n    .astype(str)\n    .isin(labeled_ids)\n].copy()\n\nprint(\n    \"Matching train studies:\",\n    labeled_train[\"StudyInstanceUID\"].nunique()\n)\n\nif labeled_train[\"StudyInstanceUID\"].nunique() != 58:\n    raise ValueError(\n        \"Expected exactly 58 fully labeled studies.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Select top candidate per study + plane\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"SELECTING PRIMARY SERIES\")\nprint(\"=\" * 70)\n\ncandidate_subset = candidates[\n    candidates[\"Anatomical_Plane\"].isin(PLANES)\n].copy()\n\ncandidate_subset = candidate_subset.sort_values(\n    [\n        \"StudyInstanceUID\",\n        \"Anatomical_Plane\",\n        \"rank_within_plane\"\n    ]\n)\n\nprimary = (\n    candidate_subset\n    .groupby(\n        [\"StudyInstanceUID\", \"Anatomical_Plane\"],\n        as_index=False\n    )\n    .first()\n)\n\nprint(\n    \"Primary series rows:\",\n    len(primary)\n)\n\n\n# ---------------------------------------------------------------\n# Check plane coverage\n# ---------------------------------------------------------------\n\nplane_counts = (\n    primary\n    .groupby(\"StudyInstanceUID\")[\"Anatomical_Plane\"]\n    .nunique()\n)\n\nprint(\n    \"Studies with all 3 primary planes:\",\n    int((plane_counts == 3).sum())\n)\n\nprint(\n    \"Studies missing at least one primary plane:\",\n    int((plane_counts < 3).sum())\n)\n\nmissing_plane_studies = plane_counts[\n    plane_counts < 3\n]\n\nif len(missing_plane_studies) > 0:\n\n    print(\"\\nWARNING: Missing-plane studies:\")\n\n    for study_id in missing_plane_studies.index:\n        present = set(\n            primary.loc[\n                primary[\"StudyInstanceUID\"] == study_id,\n                \"Anatomical_Plane\"\n            ]\n        )\n\n        missing = set(PLANES) - present\n\n        print(\n            study_id,\n            \"missing:\",\n            sorted(missing)\n        )\n\n    raise ValueError(\n        \"Every labeled study must have all 3 primary planes.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Pivot primary series into one study-level row\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BUILDING STUDY-LEVEL SERIES TABLE\")\nprint(\"=\" * 70)\n\nmodel_table = labeled_train[\n    [\"StudyInstanceUID\"] + TARGETS\n].copy()\n\nmodel_table[\"StudyInstanceUID\"] = (\n    model_table[\"StudyInstanceUID\"].astype(str)\n)\n\nfor plane in PLANES:\n\n    plane_df = primary[\n        primary[\"Anatomical_Plane\"] == plane\n    ][\n        [\n            \"StudyInstanceUID\",\n            \"SeriesInstanceUID\",\n            \"slice_count\",\n            \"Rows\",\n            \"Columns\",\n            \"PixelSpacingX\",\n            \"PixelSpacingY\",\n            \"SliceThickness\",\n            \"SpacingBetweenSlices\",\n            \"SeriesDescription\"\n        ]\n    ].copy()\n\n    plane_df = plane_df.rename(\n        columns={\n            \"SeriesInstanceUID\":\n                f\"{plane}_SeriesInstanceUID\",\n            \"slice_count\":\n                f\"{plane}_slice_count\",\n            \"Rows\":\n                f\"{plane}_Rows\",\n            \"Columns\":\n                f\"{plane}_Columns\",\n            \"PixelSpacingX\":\n                f\"{plane}_PixelSpacingX\",\n            \"PixelSpacingY\":\n                f\"{plane}_PixelSpacingY\",\n            \"SliceThickness\":\n                f\"{plane}_SliceThickness\",\n            \"SpacingBetweenSlices\":\n                f\"{plane}_SpacingBetweenSlices\",\n            \"SeriesDescription\":\n                f\"{plane}_SeriesDescription\"\n        }\n    )\n\n    model_table = model_table.merge(\n        plane_df,\n        on=\"StudyInstanceUID\",\n        how=\"left\",\n        validate=\"one_to_one\"\n    )\n\n\n# ---------------------------------------------------------------\n# Final study-level integrity checks\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"STUDY-LEVEL INTEGRITY CHECK\")\nprint(\"=\" * 70)\n\nprint(\n    \"Modeling studies:\",\n    model_table[\"StudyInstanceUID\"].nunique()\n)\n\nprint(\n    \"Modeling rows:\",\n    len(model_table)\n)\n\nif len(model_table) != 58:\n    raise ValueError(\n        \"Modeling table must contain exactly 58 studies.\"\n    )\n\nif model_table[\"StudyInstanceUID\"].duplicated().any():\n    raise ValueError(\n        \"Duplicate StudyInstanceUID detected.\"\n    )\n\n\n# Check primary series IDs\nfor plane in PLANES:\n\n    col = f\"{plane}_SeriesInstanceUID\"\n\n    missing = model_table[col].isna().sum()\n\n    unique = model_table[col].nunique()\n\n    print(\n        f\"{plane}: \"\n        f\"missing={missing}, \"\n        f\"unique_series={unique}\"\n    )\n\n    if missing != 0:\n        raise ValueError(\n            f\"Missing primary {plane} series.\"\n        )\n\n\n# ---------------------------------------------------------------\n# Verify labels are binary and complete\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TARGET VALIDATION\")\nprint(\"=\" * 70)\n\nfor target in TARGETS:\n\n    values = sorted(\n        model_table[target]\n        .dropna()\n        .unique()\n        .tolist()\n    )\n\n    print(\n        f\"{target:20s}: {values}\"\n    )\n\n    if set(values) != {0, 1}:\n        raise ValueError(\n            f\"Unexpected values in target {target}: {values}\"\n        )\n\n    if model_table[target].isna().any():\n        raise ValueError(\n            f\"Missing labels in target {target}.\"\n        )\n\n\n# ---------------------------------------------------------------\n# Create deterministic multilabel-aware folds\n# ---------------------------------------------------------------\n#\n# We use a greedy iterative assignment.\n#\n# Goal:\n#   - keep each target's positive/negative counts reasonably\n#     balanced across 3 folds\n#   - keep each fold approximately equal in size\n#   - preserve complete studies as the split unit\n#\n# This is intentionally study-level.\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CREATING 3-FOLD STUDY-LEVEL SPLIT\")\nprint(\"=\" * 70)\n\nrng = np.random.RandomState(SEED)\n\nwork = model_table[\n    [\"StudyInstanceUID\"] + TARGETS\n].copy()\n\nwork = work.sample(\n    frac=1.0,\n    random_state=SEED\n).reset_index(drop=True)\n\nY = work[TARGETS].astype(int).values\n\nn_samples = len(work)\n\n# Desired number of samples per fold\nbase_size = n_samples // N_FOLDS\nremainder = n_samples % N_FOLDS\n\ndesired_sizes = np.array([\n    base_size + (1 if i < remainder else 0)\n    for i in range(N_FOLDS)\n])\n\n# Desired positive counts per target per fold\ntarget_totals = Y.sum(axis=0)\n\ndesired_positive = (\n    target_totals / N_FOLDS\n)\n\nfold_indices = [\n    [] for _ in range(N_FOLDS)\n]\n\nfold_positive = np.zeros(\n    (N_FOLDS, len(TARGETS)),\n    dtype=float\n)\n\nfold_sizes = np.zeros(\n    N_FOLDS,\n    dtype=int\n)\n\n# Process studies with more abnormalities first.\n# This helps distribute rare multilabel signatures.\nseverity = Y.sum(axis=1)\n\norder = np.argsort(\n    -severity +\n    rng.uniform(\n        0,\n        0.01,\n        size=n_samples\n    )\n)\n\nfor idx in order:\n\n    sample_labels = Y[idx]\n\n    best_fold = None\n    best_score = None\n\n    for fold in range(N_FOLDS):\n\n        if (\n            fold_sizes[fold]\n            >= desired_sizes[fold]\n        ):\n            continue\n\n        new_positive = (\n            fold_positive[fold]\n            + sample_labels\n        )\n\n        # Positive-count deviation\n        positive_error = np.sum(\n            (\n                new_positive\n                - desired_positive\n            ) ** 2\n        )\n\n        # Fold-size penalty\n        size_error = (\n            (\n                fold_sizes[fold] + 1\n            )\n            - desired_sizes[fold]\n        ) ** 2\n\n        # Small penalty for already-large folds\n        load_penalty = (\n            fold_sizes[fold]\n            / max(1, desired_sizes.max())\n        )\n\n        score = (\n            positive_error\n            + 0.25 * size_error\n            + 0.05 * load_penalty\n        )\n\n        if (\n            best_score is None\n            or score < best_score\n        ):\n            best_score = score\n            best_fold = fold\n\n    fold_indices[best_fold].append(idx)\n\n    fold_positive[best_fold] += sample_labels\n    fold_sizes[best_fold] += 1\n\n\n# ---------------------------------------------------------------\n# Convert assignments back to StudyInstanceUID\n# ---------------------------------------------------------------\n\nfold_assignment = {}\n\nfor fold in range(N_FOLDS):\n\n    for idx in fold_indices[fold]:\n\n        study_id = work.loc[\n            idx,\n            \"StudyInstanceUID\"\n        ]\n\n        fold_assignment[\n            study_id\n        ] = fold\n\n\nmodel_table[\"fold\"] = (\n    model_table[\"StudyInstanceUID\"]\n    .map(fold_assignment)\n    .astype(int)\n)\n\n\n# ---------------------------------------------------------------\n# Fold diagnostics\n# ---------------------------------------------------------------\n\nprint()\n\nfor fold in range(N_FOLDS):\n\n    fold_df = model_table[\n        model_table[\"fold\"] == fold\n    ]\n\n    print(\n        f\"Fold {fold + 1}: \"\n        f\"{len(fold_df)} studies\"\n    )\n\n    for target in TARGETS:\n\n        pos = int(\n            fold_df[target].sum()\n        )\n\n        neg = int(\n            len(fold_df) - pos\n        )\n\n        print(\n            f\"  {target:20s} \"\n            f\"pos={pos:2d} \"\n            f\"neg={neg:2d}\"\n        )\n\n    print()\n\n\n# ---------------------------------------------------------------\n# Verify every fold has both classes\n# ---------------------------------------------------------------\n\nprint(\"=\" * 70)\nprint(\"CLASS COVERAGE CHECK\")\nprint(\"=\" * 70)\n\nbad = []\n\nfor fold in range(N_FOLDS):\n\n    fold_df = model_table[\n        model_table[\"fold\"] == fold\n    ]\n\n    for target in TARGETS:\n\n        pos = int(\n            fold_df[target].sum()\n        )\n\n        neg = int(\n            len(fold_df) - pos\n        )\n\n        if pos == 0 or neg == 0:\n\n            bad.append(\n                (\n                    fold,\n                    target,\n                    pos,\n                    neg\n                )\n            )\n\nif bad:\n\n    print(\n        \"WARNING: Some target/fold combinations \"\n        \"do not contain both classes.\"\n    )\n\n    for item in bad:\n        print(item)\n\nelse:\n\n    print(\n        \"PASS: Every target has both classes \"\n        \"in every fold.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Verify no study leakage\n# ---------------------------------------------------------------\n\nfold_counts = (\n    model_table\n    .groupby(\"StudyInstanceUID\")[\"fold\"]\n    .nunique()\n)\n\nif (fold_counts != 1).any():\n\n    raise ValueError(\n        \"Study leakage detected across folds.\"\n    )\n\nprint(\n    \"PASS: Every study belongs to exactly one fold.\"\n)\n\n\n# ---------------------------------------------------------------\n# Save\n# ---------------------------------------------------------------\n\nmodel_table.to_csv(\n    OUTPUT_MODEL_TABLE,\n    index=False\n)\n\nmodel_table[\n    [\"StudyInstanceUID\", \"fold\"]\n].to_csv(\n    OUTPUT_FOLDS,\n    index=False\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 29 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved:\",\n    OUTPUT_MODEL_TABLE\n)\n\nprint(\n    \"Saved:\",\n    OUTPUT_FOLDS\n)\n\nprint()\nprint(\n    \"Modeling table shape:\",\n    model_table.shape\n)\n\nprint(\n    \"Fold distribution:\"\n)\n\nprint(\n    model_table[\"fold\"]\n    .value_counts()\n    .sort_index()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:46:48.768813Z","iopub.execute_input":"2026-08-11T08:46:48.769625Z","iopub.status.idle":"2026-08-11T08:46:48.966949Z","shell.execute_reply.started":"2026-08-11T08:46:48.769584Z","shell.execute_reply":"2026-08-11T08:46:48.965996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 29B - ROBUST MULTILABEL 3-FOLD SPLIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODEL_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nOUTPUT_MODEL_TABLE = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nOUTPUT_FOLDS = os.path.join(\n    AUDIT_DIR,\n    \"cell29_study_folds.csv\"\n)\n\nSEED = 42\nN_FOLDS = 3\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\nprint(\"=\" * 70)\nprint(\"CELL 29B - ROBUST MULTILABEL 3-FOLD SPLIT\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load existing modeling table\n# ---------------------------------------------------------------\n\nmodel_table = pd.read_csv(\n    MODEL_TABLE_PATH\n)\n\nprint(\n    \"\\nLoaded modeling table:\",\n    model_table.shape\n)\n\nif len(model_table) != 58:\n    raise ValueError(\n        f\"Expected 58 labeled studies, got {len(model_table)}\"\n    )\n\nif model_table[\"StudyInstanceUID\"].duplicated().any():\n    raise ValueError(\n        \"Duplicate StudyInstanceUID detected.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Prepare labels\n# ---------------------------------------------------------------\n\nY = (\n    model_table[TARGETS]\n    .astype(int)\n    .values\n)\n\nn_samples = len(model_table)\nn_targets = len(TARGETS)\n\nif n_samples != 58:\n    raise ValueError(\n        \"Expected exactly 58 labeled studies.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Objective function\n# ---------------------------------------------------------------\n#\n# We want:\n#\n# 1. Every fold to have both classes for every target.\n# 2. Positive counts to be close to target_total / 3.\n# 3. Fold sizes close to 58 / 3.\n#\n# Invalid folds receive a very large penalty.\n# ---------------------------------------------------------------\n\ntarget_totals = Y.sum(axis=0)\n\ndesired_positive = (\n    target_totals / N_FOLDS\n)\n\ndesired_size = (\n    n_samples / N_FOLDS\n)\n\n\ndef evaluate_assignment(\n    assignment\n):\n\n    assignment = np.asarray(\n        assignment\n    )\n\n    score = 0.0\n\n    fold_sizes = np.zeros(\n        N_FOLDS,\n        dtype=int\n    )\n\n    fold_positive = np.zeros(\n        (N_FOLDS, n_targets),\n        dtype=int\n    )\n\n    for i in range(n_samples):\n\n        fold = assignment[i]\n\n        fold_sizes[fold] += 1\n        fold_positive[fold] += Y[i]\n\n    # -----------------------------------------------------------\n    # Fold-size balance\n    # -----------------------------------------------------------\n\n    score += np.sum(\n        (\n            fold_sizes\n            - desired_size\n        ) ** 2\n    ) * 2.0\n\n    # -----------------------------------------------------------\n    # Positive-count balance\n    # -----------------------------------------------------------\n\n    for fold in range(N_FOLDS):\n\n        diff = (\n            fold_positive[fold]\n            - desired_positive\n        )\n\n        # Normalize by target frequency so rare targets\n        # receive appropriate importance.\n        normalized = (\n            diff ** 2\n            / np.maximum(\n                desired_positive,\n                1.0\n            )\n        )\n\n        score += np.sum(\n            normalized\n        )\n\n    # -----------------------------------------------------------\n    # Hard class-coverage constraints\n    # -----------------------------------------------------------\n\n    invalid_count = 0\n\n    for fold in range(N_FOLDS):\n\n        positives = fold_positive[fold]\n\n        fold_size = fold_sizes[fold]\n\n        negatives = (\n            fold_size - positives\n        )\n\n        invalid_count += np.sum(\n            positives == 0\n        )\n\n        invalid_count += np.sum(\n            negatives == 0\n        )\n\n    # Huge penalty for invalid target/fold combinations.\n    score += (\n        invalid_count * 10000.0\n    )\n\n    return (\n        score,\n        fold_sizes,\n        fold_positive\n    )\n\n\n# ---------------------------------------------------------------\n# Randomized local-search optimization\n# ---------------------------------------------------------------\n\nrng = np.random.RandomState(\n    SEED\n)\n\nbest_assignment = None\nbest_score = np.inf\nbest_details = None\n\nN_RESTARTS = 300\nMAX_ITERATIONS = 5000\n\nprint()\nprint(\n    \"Searching for a valid multilabel split...\"\n)\n\nfor restart in range(N_RESTARTS):\n\n    # Start with balanced fold assignment.\n    assignment = np.arange(\n        n_samples\n    ) % N_FOLDS\n\n    rng.shuffle(\n        assignment\n    )\n\n    current_score, _, _ = evaluate_assignment(\n        assignment\n    )\n\n    for iteration in range(\n        MAX_ITERATIONS\n    ):\n\n        improved = False\n\n        # Randomly inspect study pairs.\n        pairs = rng.randint(\n            0,\n            n_samples,\n            size=(80, 2)\n        )\n\n        for i, j in pairs:\n\n            if i == j:\n                continue\n\n            if assignment[i] == assignment[j]:\n                continue\n\n            old_i = assignment[i]\n            old_j = assignment[j]\n\n            assignment[i] = old_j\n            assignment[j] = old_i\n\n            new_score, _, _ = evaluate_assignment(\n                assignment\n            )\n\n            if new_score < current_score:\n\n                current_score = new_score\n                improved = True\n\n            else:\n\n                assignment[i] = old_i\n                assignment[j] = old_j\n\n        if not improved:\n            break\n\n    final_score, fold_sizes, fold_positive = (\n        evaluate_assignment(\n            assignment\n        )\n    )\n\n    if final_score < best_score:\n\n        best_score = final_score\n        best_assignment = assignment.copy()\n\n        best_details = (\n            fold_sizes.copy(),\n            fold_positive.copy()\n        )\n\n    # Stop immediately once a valid and reasonably balanced\n    # solution is found.\n    if (\n        np.all(\n            best_details[1] > 0\n        )\n        and np.all(\n            best_details[1]\n            < best_details[0][:, None]\n        )\n    ):\n\n        # We have a valid solution.\n        if restart >= 10:\n            break\n\n\n# ---------------------------------------------------------------\n# Safety check\n# ---------------------------------------------------------------\n\nif best_assignment is None:\n    raise RuntimeError(\n        \"Could not generate a fold assignment.\"\n    )\n\nfold_sizes, fold_positive = best_details\n\nprint()\nprint(\n    \"Best objective score:\",\n    best_score\n)\n\nprint(\n    \"Fold sizes:\",\n    fold_sizes.tolist()\n)\n\n\n# ---------------------------------------------------------------\n# Validate class coverage\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL FOLD DIAGNOSTICS\")\nprint(\"=\" * 70)\n\ninvalid = []\n\nfor fold in range(N_FOLDS):\n\n    positives = fold_positive[fold]\n\n    negatives = (\n        fold_sizes[fold]\n        - positives\n    )\n\n    print()\n    print(\n        f\"Fold {fold + 1}: \"\n        f\"{fold_sizes[fold]} studies\"\n    )\n\n    for t, target in enumerate(TARGETS):\n\n        pos = int(\n            positives[t]\n        )\n\n        neg = int(\n            negatives[t]\n        )\n\n        print(\n            f\"  {target:20s} \"\n            f\"pos={pos:2d} \"\n            f\"neg={neg:2d}\"\n        )\n\n        if pos == 0 or neg == 0:\n\n            invalid.append(\n                (\n                    fold,\n                    target,\n                    pos,\n                    neg\n                )\n            )\n\n\n# ---------------------------------------------------------------\n# Apply fold assignment\n# ---------------------------------------------------------------\n\nmodel_table[\"fold\"] = (\n    best_assignment\n)\n\n\n# ---------------------------------------------------------------\n# Verify no study leakage\n# ---------------------------------------------------------------\n\nfold_counts = (\n    model_table\n    .groupby(\n        \"StudyInstanceUID\"\n    )[\"fold\"]\n    .nunique()\n)\n\nif (fold_counts != 1).any():\n\n    raise ValueError(\n        \"A study appears in multiple folds.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Final verdict\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CLASS COVERAGE VERDICT\")\nprint(\"=\" * 70)\n\nif invalid:\n\n    print(\n        \"FAIL: Invalid target/fold combinations:\"\n    )\n\n    for item in invalid:\n        print(item)\n\n    raise RuntimeError(\n        \"Could not construct a valid 3-fold split.\"\n    )\n\nelse:\n\n    print(\n        \"PASS: Every target has both positive \"\n        \"and negative samples in every fold.\"\n    )\n\nprint(\n    \"PASS: Every study belongs to exactly one fold.\"\n)\n\n\n# ---------------------------------------------------------------\n# Save\n# ---------------------------------------------------------------\n\nmodel_table.to_csv(\n    OUTPUT_MODEL_TABLE,\n    index=False\n)\n\nmodel_table[\n    [\n        \"StudyInstanceUID\",\n        \"fold\"\n    ]\n].to_csv(\n    OUTPUT_FOLDS,\n    index=False\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 29B COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved:\",\n    OUTPUT_MODEL_TABLE\n)\n\nprint(\n    \"Saved:\",\n    OUTPUT_FOLDS\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:47:40.660563Z","iopub.execute_input":"2026-08-11T08:47:40.660949Z","iopub.status.idle":"2026-08-11T08:47:41.196402Z","shell.execute_reply.started":"2026-08-11T08:47:40.660908Z","shell.execute_reply":"2026-08-11T08:47:41.195463Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 30 - ROBUST DICOM SERIES LOADER\n# ================================================================\n\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\n\nfrom PIL import Image\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nTRAIN_SERIES_DIR = (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/train_series\"\n)\n\nMODEL_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nIMAGE_SIZE = 224\nNUM_SLICES = 5\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\nSEED = 42\n\nprint(\"=\" * 70)\nprint(\"CELL 30 - DICOM MODEL INPUT LOADER\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load modeling table\n# ---------------------------------------------------------------\n\nmodel_table = pd.read_csv(\n    MODEL_TABLE_PATH\n)\n\nprint(\n    \"\\nModeling studies:\",\n    len(model_table)\n)\n\nprint(\n    \"Image size:\",\n    IMAGE_SIZE\n)\n\nprint(\n    \"Slices per series:\",\n    NUM_SLICES\n)\n\nprint(\n    \"Planes:\",\n    PLANES\n)\n\n\n# ---------------------------------------------------------------\n# Utility: find DICOM files\n# ---------------------------------------------------------------\n\ndef get_dicom_files(series_path):\n    \"\"\"\n    Return DICOM files from a series directory.\n    \"\"\"\n\n    if not os.path.isdir(series_path):\n        return []\n\n    files = []\n\n    for filename in os.listdir(series_path):\n\n        path = os.path.join(\n            series_path,\n            filename\n        )\n\n        if os.path.isfile(path):\n            files.append(path)\n\n    return files\n\n\n# ---------------------------------------------------------------\n# Utility: calculate slice position\n# ---------------------------------------------------------------\n\ndef get_slice_position(ds):\n    \"\"\"\n    Calculate the physical position of a DICOM slice.\n\n    Uses ImagePositionPatient and\n    ImageOrientationPatient when available.\n    \"\"\"\n\n    position = getattr(\n        ds,\n        \"ImagePositionPatient\",\n        None\n    )\n\n    orientation = getattr(\n        ds,\n        \"ImageOrientationPatient\",\n        None\n    )\n\n    if (\n        position is not None\n        and orientation is not None\n        and len(position) >= 3\n        and len(orientation) >= 6\n    ):\n\n        try:\n\n            row_cosines = np.asarray(\n                orientation[:3],\n                dtype=np.float64\n            )\n\n            col_cosines = np.asarray(\n                orientation[3:6],\n                dtype=np.float64\n            )\n\n            normal = np.cross(\n                row_cosines,\n                col_cosines\n            )\n\n            position = np.asarray(\n                position,\n                dtype=np.float64\n            )\n\n            return float(\n                np.dot(\n                    position,\n                    normal\n                )\n            )\n\n        except Exception:\n            pass\n\n    # Fallback to InstanceNumber\n    instance = getattr(\n        ds,\n        \"InstanceNumber\",\n        None\n    )\n\n    if instance is not None:\n\n        try:\n            return float(instance)\n        except Exception:\n            pass\n\n    return 0.0\n\n\n# ---------------------------------------------------------------\n# Utility: sort DICOM files\n# ---------------------------------------------------------------\n\ndef sort_dicom_files(dicom_files):\n    \"\"\"\n    Sort DICOM files using physical slice position.\n    \"\"\"\n\n    records = []\n\n    for path in dicom_files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                path,\n                stop_before_pixels=True,\n                force=True\n            )\n\n            position = get_slice_position(\n                ds\n            )\n\n            records.append(\n                (\n                    position,\n                    path\n                )\n            )\n\n        except Exception:\n            continue\n\n    records.sort(\n        key=lambda x: x[0]\n    )\n\n    return [\n        path\n        for _, path in records\n    ]\n\n\n# ---------------------------------------------------------------\n# Utility: robust intensity normalization\n# ---------------------------------------------------------------\n\ndef normalize_dicom_image(\n    pixel_array\n):\n    \"\"\"\n    Robust percentile-based normalization.\n\n    The purpose is to reduce sensitivity to\n    different MRI intensity scales.\n    \"\"\"\n\n    image = np.asarray(\n        pixel_array,\n        dtype=np.float32\n    )\n\n    image = np.nan_to_num(\n        image,\n        nan=0.0,\n        posinf=0.0,\n        neginf=0.0\n    )\n\n    # Remove extreme values.\n    low = np.percentile(\n        image,\n        1\n    )\n\n    high = np.percentile(\n        image,\n        99\n    )\n\n    if high <= low:\n\n        low = float(\n            image.min()\n        )\n\n        high = float(\n            image.max()\n        )\n\n    if high <= low:\n\n        return np.zeros_like(\n            image,\n            dtype=np.float32\n        )\n\n    image = np.clip(\n        image,\n        low,\n        high\n    )\n\n    image = (\n        image - low\n    ) / (\n        high - low\n    )\n\n    return image.astype(\n        np.float32\n    )\n\n\n# ---------------------------------------------------------------\n# Utility: resize image\n# ---------------------------------------------------------------\n\ndef resize_image(\n    image,\n    size=IMAGE_SIZE\n):\n    \"\"\"\n    Resize normalized image to model resolution.\n    \"\"\"\n\n    image_uint8 = (\n        np.clip(\n            image,\n            0.0,\n            1.0\n        )\n        * 255.0\n    ).astype(\n        np.uint8\n    )\n\n    pil_image = Image.fromarray(\n        image_uint8\n    )\n\n    pil_image = pil_image.resize(\n        (size, size),\n        resample=Image.BILINEAR\n    )\n\n    output = np.asarray(\n        pil_image,\n        dtype=np.float32\n    ) / 255.0\n\n    return output\n\n\n# ---------------------------------------------------------------\n# Select representative slice indices\n# ---------------------------------------------------------------\n\ndef select_slice_indices(\n    n_slices,\n    num_slices=NUM_SLICES\n):\n    \"\"\"\n    Select approximately uniformly spaced slices.\n    \"\"\"\n\n    if n_slices <= 0:\n        return []\n\n    if n_slices <= num_slices:\n\n        return list(\n            range(n_slices)\n        )\n\n    indices = np.linspace(\n        0,\n        n_slices - 1,\n        num=num_slices\n    )\n\n    indices = np.round(\n        indices\n    ).astype(int)\n\n    return list(\n        np.unique(indices)\n    )\n\n\n# ---------------------------------------------------------------\n# Load a single series\n# ---------------------------------------------------------------\n\ndef load_series(\n    study_id,\n    series_id,\n    num_slices=NUM_SLICES,\n    image_size=IMAGE_SIZE\n):\n    \"\"\"\n    Load representative slices from one DICOM series.\n\n    Returns:\n        tensor-like numpy array:\n        [N, H, W]\n    \"\"\"\n\n    series_path = os.path.join(\n        TRAIN_SERIES_DIR,\n        str(study_id),\n        str(series_id)\n    )\n\n    files = get_dicom_files(\n        series_path\n    )\n\n    if len(files) == 0:\n\n        raise FileNotFoundError(\n            f\"No DICOM files found: \"\n            f\"{series_path}\"\n        )\n\n    sorted_files = sort_dicom_files(\n        files\n    )\n\n    if len(sorted_files) == 0:\n\n        raise RuntimeError(\n            f\"Could not sort/read DICOM files: \"\n            f\"{series_path}\"\n        )\n\n    selected_indices = (\n        select_slice_indices(\n            len(sorted_files),\n            num_slices\n        )\n    )\n\n    images = []\n\n    for index in selected_indices:\n\n        path = sorted_files[index]\n\n        try:\n\n            ds = pydicom.dcmread(\n                path,\n                force=True\n            )\n\n            if not hasattr(\n                ds,\n                \"PixelData\"\n            ):\n\n                continue\n\n            image = ds.pixel_array\n\n            # Apply DICOM rescaling when available.\n            slope = float(\n                getattr(\n                    ds,\n                    \"RescaleSlope\",\n                    1.0\n                )\n            )\n\n            intercept = float(\n                getattr(\n                    ds,\n                    \"RescaleIntercept\",\n                    0.0\n                )\n            )\n\n            image = (\n                image.astype(\n                    np.float32\n                )\n                * slope\n                + intercept\n            )\n\n            # Handle MONOCHROME1 safely.\n            photometric = getattr(\n                ds,\n                \"PhotometricInterpretation\",\n                \"MONOCHROME2\"\n            )\n\n            if photometric == \"MONOCHROME1\":\n\n                image = (\n                    image.max()\n                    - image\n                )\n\n            image = normalize_dicom_image(\n                image\n            )\n\n            image = resize_image(\n                image,\n                image_size\n            )\n\n            images.append(\n                image\n            )\n\n        except Exception as e:\n\n            print(\n                \"WARNING: Failed to read:\",\n                path,\n                \"|\",\n                type(e).__name__\n            )\n\n    if len(images) == 0:\n\n        raise RuntimeError(\n            f\"No usable slices found for \"\n            f\"{study_id} / {series_id}\"\n        )\n\n    return np.stack(\n        images,\n        axis=0\n    ).astype(\n        np.float32\n    )\n\n\n# ---------------------------------------------------------------\n# Test loader on one study\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TESTING LOADER\")\nprint(\"=\" * 70)\n\ntest_row = model_table.iloc[0]\n\nstudy_id = str(\n    test_row[\"StudyInstanceUID\"]\n)\n\nprint(\n    \"Test study:\",\n    study_id\n)\n\nloaded_series = {}\n\nfor plane in PLANES:\n\n    series_id = str(\n        test_row[\n            f\"{plane}_SeriesInstanceUID\"\n        ]\n    )\n\n    print()\n    print(\n        f\"Loading {plane}:\"\n    )\n\n    print(\n        \"Series:\",\n        series_id\n    )\n\n    volume = load_series(\n        study_id,\n        series_id\n    )\n\n    loaded_series[\n        plane\n    ] = volume\n\n    print(\n        \"Loaded shape:\",\n        volume.shape\n    )\n\n    print(\n        \"Min:\",\n        float(volume.min())\n    )\n\n    print(\n        \"Max:\",\n        float(volume.max())\n    )\n\n    print(\n        \"Mean:\",\n        float(volume.mean())\n    )\n\n    print(\n        \"Std:\",\n        float(volume.std())\n    )\n\n\n# ---------------------------------------------------------------\n# Combine all planes\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"COMBINING PRIMARY PLANES\")\nprint(\"=\" * 70)\n\ncombined = np.concatenate(\n    [\n        loaded_series[\"Sagittal\"],\n        loaded_series[\"Coronal\"],\n        loaded_series[\"Axial\"]\n    ],\n    axis=0\n)\n\nprint(\n    \"Combined shape:\",\n    combined.shape\n)\n\nprint(\n    \"Expected approximately:\",\n    NUM_SLICES * len(PLANES),\n    \"slices\"\n)\n\nprint(\n    \"Tensor memory:\",\n    round(\n        combined.nbytes\n        / (1024 ** 2),\n        4\n    ),\n    \"MB\"\n)\n\n\n# ---------------------------------------------------------------\n# Convert to PyTorch tensor\n# ---------------------------------------------------------------\n\ntensor = torch.from_numpy(\n    combined\n)\n\nprint(\n    \"PyTorch tensor shape:\",\n    tuple(tensor.shape)\n)\n\nprint(\n    \"PyTorch dtype:\",\n    tensor.dtype\n)\n\n\n# ---------------------------------------------------------------\n# Cleanup\n# ---------------------------------------------------------------\n\ndel loaded_series\ndel combined\ndel tensor\n\ngc.collect()\n\n\n# ---------------------------------------------------------------\n# Final\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 30 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"DICOM loader test completed successfully.\"\n)\n\nprint(\n    \"No full dataset was loaded into memory.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:49:06.321099Z","iopub.execute_input":"2026-08-11T08:49:06.321393Z","iopub.status.idle":"2026-08-11T08:49:06.917031Z","shell.execute_reply.started":"2026-08-11T08:49:06.321359Z","shell.execute_reply":"2026-08-11T08:49:06.916223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 31 - VISUALIZE EXACT MODEL INPUTS\n# ================================================================\n\nimport os\nimport math\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODEL_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nOUTPUT_IMAGE = os.path.join(\n    AUDIT_DIR,\n    \"cell31_model_input_contact_sheet.png\"\n)\n\nIMAGE_SIZE = 224\nNUM_SLICES = 5\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\nprint(\"=\" * 70)\nprint(\"CELL 31 - MODEL INPUT VISUALIZATION\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load modeling table\n# ---------------------------------------------------------------\n\nmodel_table = pd.read_csv(\n    MODEL_TABLE_PATH\n)\n\ntest_row = model_table.iloc[0]\n\nstudy_id = str(\n    test_row[\"StudyInstanceUID\"]\n)\n\nprint(\"\\nStudy:\")\nprint(study_id)\n\n\n# ---------------------------------------------------------------\n# Load series using Cell 30 loader\n# ---------------------------------------------------------------\n\nloaded = {}\n\nfor plane in PLANES:\n\n    series_id = str(\n        test_row[\n            f\"{plane}_SeriesInstanceUID\"\n        ]\n    )\n\n    print(\n        f\"\\nLoading {plane}:\"\n    )\n\n    volume = load_series(\n        study_id,\n        series_id,\n        num_slices=NUM_SLICES,\n        image_size=IMAGE_SIZE\n    )\n\n    loaded[plane] = volume\n\n    print(\n        \"Shape:\",\n        volume.shape\n    )\n\n\n# ---------------------------------------------------------------\n# Create contact sheet\n# ---------------------------------------------------------------\n\nfig, axes = plt.subplots(\n    len(PLANES),\n    NUM_SLICES,\n    figsize=(15, 9)\n)\n\nif len(PLANES) == 1:\n    axes = np.expand_dims(\n        axes,\n        axis=0\n    )\n\nfor row_idx, plane in enumerate(PLANES):\n\n    volume = loaded[plane]\n\n    for col_idx in range(NUM_SLICES):\n\n        ax = axes[\n            row_idx,\n            col_idx\n        ]\n\n        ax.imshow(\n            volume[col_idx],\n            cmap=\"gray\",\n            vmin=0,\n            vmax=1\n        )\n\n        ax.axis(\"off\")\n\n        if col_idx == 0:\n\n            ax.set_title(\n                plane,\n                fontsize=12\n            )\n\n        else:\n\n            ax.set_title(\n                f\"Slice {col_idx + 1}\",\n                fontsize=10\n            )\n\n\nfig.suptitle(\n    \"Exact Model Inputs - 5 Slices per Primary Plane\",\n    fontsize=15\n)\n\nplt.tight_layout()\n\nplt.savefig(\n    OUTPUT_IMAGE,\n    dpi=150,\n    bbox_inches=\"tight\"\n)\n\nplt.show()\n\nplt.close(fig)\n\n\n# ---------------------------------------------------------------\n# Print numerical summary\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"INPUT SUMMARY\")\nprint(\"=\" * 70)\n\nfor plane in PLANES:\n\n    volume = loaded[plane]\n\n    print(\n        f\"{plane:10s} | \"\n        f\"shape={volume.shape} | \"\n        f\"min={volume.min():.4f} | \"\n        f\"max={volume.max():.4f} | \"\n        f\"mean={volume.mean():.4f} | \"\n        f\"std={volume.std():.4f}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Save summary\n# ---------------------------------------------------------------\n\nsummary_rows = []\n\nfor plane in PLANES:\n\n    volume = loaded[plane]\n\n    summary_rows.append({\n        \"StudyInstanceUID\": study_id,\n        \"Plane\": plane,\n        \"num_slices\": volume.shape[0],\n        \"height\": volume.shape[1],\n        \"width\": volume.shape[2],\n        \"min\": float(volume.min()),\n        \"max\": float(volume.max()),\n        \"mean\": float(volume.mean()),\n        \"std\": float(volume.std())\n    })\n\nsummary_df = pd.DataFrame(\n    summary_rows\n)\n\nsummary_path = os.path.join(\n    AUDIT_DIR,\n    \"cell31_model_input_summary.csv\"\n)\n\nsummary_df.to_csv(\n    summary_path,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Final\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 31 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved visualization:\",\n    OUTPUT_IMAGE\n)\n\nprint(\n    \"Saved summary:\",\n    summary_path\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:50:14.870825Z","iopub.execute_input":"2026-08-11T08:50:14.871412Z","iopub.status.idle":"2026-08-11T08:50:18.166582Z","shell.execute_reply.started":"2026-08-11T08:50:14.871381Z","shell.execute_reply":"2026-08-11T08:50:18.165696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 32 - SLICE SAMPLING STRATEGY AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODEL_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nIMAGE_SIZE = 224\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\n# Number of studies to visually inspect.\nNUM_STUDIES = 3\n\n# Number of slices to select.\nNUM_SLICES = 5\n\nprint(\"=\" * 70)\nprint(\"CELL 32 - SLICE SAMPLING STRATEGY AUDIT\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Sampling strategies\n# ---------------------------------------------------------------\n\ndef uniform_indices(\n    n,\n    k=5\n):\n    \"\"\"\n    Current strategy used in Cell 30.\n    \"\"\"\n\n    if n <= k:\n        return list(range(n))\n\n    indices = np.linspace(\n        0,\n        n - 1,\n        k\n    )\n\n    return list(\n        np.unique(\n            np.round(indices).astype(int)\n        )\n    )\n\n\ndef center_indices(\n    n,\n    k=5,\n    center_fraction=0.50\n):\n    \"\"\"\n    Select k slices from the central portion\n    of the series.\n    \"\"\"\n\n    if n <= k:\n        return list(range(n))\n\n    width = int(\n        round(\n            n * center_fraction\n        )\n    )\n\n    width = max(\n        width,\n        k\n    )\n\n    start = max(\n        0,\n        (n - width) // 2\n    )\n\n    end = min(\n        n - 1,\n        start + width - 1\n    )\n\n    indices = np.linspace(\n        start,\n        end,\n        k\n    )\n\n    return list(\n        np.unique(\n            np.round(indices).astype(int)\n        )\n    )\n\n\ndef hybrid_indices(\n    n,\n    k=5,\n    margin_fraction=0.10\n):\n    \"\"\"\n    Avoid extreme ends while retaining a wider\n    portion of the anatomical coverage.\n    \"\"\"\n\n    if n <= k:\n        return list(range(n))\n\n    margin = int(\n        round(\n            n * margin_fraction\n        )\n    )\n\n    start = margin\n    end = n - 1 - margin\n\n    if end <= start:\n        return uniform_indices(\n            n,\n            k\n        )\n\n    indices = np.linspace(\n        start,\n        end,\n        k\n    )\n\n    return list(\n        np.unique(\n            np.round(indices).astype(int)\n        )\n    )\n\n\n# ---------------------------------------------------------------\n# DICOM file discovery\n# ---------------------------------------------------------------\n\ndef get_sorted_dicom_files(\n    study_id,\n    series_id\n):\n    \"\"\"\n    Return DICOM files sorted using physical\n    slice position.\n    \"\"\"\n\n    series_path = os.path.join(\n        TRAIN_SERIES_DIR,\n        str(study_id),\n        str(series_id)\n    )\n\n    files = []\n\n    for filename in os.listdir(\n        series_path\n    ):\n\n        path = os.path.join(\n            series_path,\n            filename\n        )\n\n        if os.path.isfile(path):\n            files.append(path)\n\n    records = []\n\n    for path in files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                path,\n                stop_before_pixels=True,\n                force=True\n            )\n\n            position = get_slice_position(\n                ds\n            )\n\n            records.append(\n                (\n                    position,\n                    path\n                )\n            )\n\n        except Exception:\n            continue\n\n    records.sort(\n        key=lambda x: x[0]\n    )\n\n    return [\n        path\n        for _, path in records\n    ]\n\n\n# ---------------------------------------------------------------\n# Load selected slices\n# ---------------------------------------------------------------\n\ndef load_selected_images(\n    study_id,\n    series_id,\n    indices\n):\n    \"\"\"\n    Load and normalize selected DICOM slices.\n    \"\"\"\n\n    files = get_sorted_dicom_files(\n        study_id,\n        series_id\n    )\n\n    images = []\n\n    for index in indices:\n\n        if index >= len(files):\n            continue\n\n        try:\n\n            ds = pydicom.dcmread(\n                files[index],\n                force=True\n            )\n\n            image = ds.pixel_array.astype(\n                np.float32\n            )\n\n            slope = float(\n                getattr(\n                    ds,\n                    \"RescaleSlope\",\n                    1.0\n                )\n            )\n\n            intercept = float(\n                getattr(\n                    ds,\n                    \"RescaleIntercept\",\n                    0.0\n                )\n            )\n\n            image = (\n                image * slope\n                + intercept\n            )\n\n            photometric = getattr(\n                ds,\n                \"PhotometricInterpretation\",\n                \"MONOCHROME2\"\n            )\n\n            if photometric == \"MONOCHROME1\":\n\n                image = (\n                    image.max()\n                    - image\n                )\n\n            image = normalize_dicom_image(\n                image\n            )\n\n            image = resize_image(\n                image,\n                IMAGE_SIZE\n            )\n\n            images.append(\n                image\n            )\n\n        except Exception:\n            continue\n\n    return images, len(files)\n\n\n# ---------------------------------------------------------------\n# Load modeling table\n# ---------------------------------------------------------------\n\nmodel_table = pd.read_csv(\n    MODEL_TABLE_PATH\n)\n\nprint(\n    \"\\nTotal modeling studies:\",\n    len(model_table)\n)\n\nstudy_rows = model_table.iloc[\n    :NUM_STUDIES\n]\n\n\n# ---------------------------------------------------------------\n# Compare strategies\n# ---------------------------------------------------------------\n\nstrategies = {\n    \"Uniform\": uniform_indices,\n    \"Center\": center_indices,\n    \"Hybrid\": hybrid_indices\n}\n\nresults = []\n\n\nfor study_number, (_, row) in enumerate(\n    study_rows.iterrows(),\n    start=1\n):\n\n    study_id = str(\n        row[\"StudyInstanceUID\"]\n    )\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        f\"Study {study_number}:\"\n    )\n\n    print(\n        study_id\n    )\n\n    for plane in PLANES:\n\n        series_id = str(\n            row[\n                f\"{plane}_SeriesInstanceUID\"\n            ]\n        )\n\n        files = get_sorted_dicom_files(\n            study_id,\n            series_id\n        )\n\n        n = len(files)\n\n        print()\n        print(\n            f\"{plane} | total slices = {n}\"\n        )\n\n        for strategy_name, strategy_function in strategies.items():\n\n            indices = strategy_function(\n                n,\n                NUM_SLICES\n            )\n\n            results.append({\n                \"StudyInstanceUID\": study_id,\n                \"Plane\": plane,\n                \"Strategy\": strategy_name,\n                \"TotalSlices\": n,\n                \"SelectedIndices\": \",\".join(\n                    map(\n                        str,\n                        indices\n                    )\n                )\n            })\n\n            print(\n                f\"{strategy_name:8s}:\",\n                indices\n            )\n\n\n# ---------------------------------------------------------------\n# Create visual comparison\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CREATING VISUAL COMPARISON\")\nprint(\"=\" * 70)\n\n\n# Use first study and all three planes.\nrow = study_rows.iloc[0]\n\nstudy_id = str(\n    row[\"StudyInstanceUID\"]\n)\n\n\nfig, axes = plt.subplots(\n    len(PLANES) * 3,\n    NUM_SLICES,\n    figsize=(15, 24)\n)\n\n\nfor plane_idx, plane in enumerate(\n    PLANES\n):\n\n    series_id = str(\n        row[\n            f\"{plane}_SeriesInstanceUID\"\n        ]\n    )\n\n    files = get_sorted_dicom_files(\n        study_id,\n        series_id\n    )\n\n    n = len(files)\n\n    for strategy_idx, (\n        strategy_name,\n        strategy_function\n    ) in enumerate(\n        strategies.items()\n    ):\n\n        row_idx = (\n            plane_idx * 3\n            + strategy_idx\n        )\n\n        indices = strategy_function(\n            n,\n            NUM_SLICES\n        )\n\n        images, _ = load_selected_images(\n            study_id,\n            series_id,\n            indices\n        )\n\n        for col_idx in range(\n            NUM_SLICES\n        ):\n\n            ax = axes[\n                row_idx,\n                col_idx\n            ]\n\n            ax.axis(\"off\")\n\n            if col_idx < len(images):\n\n                ax.imshow(\n                    images[col_idx],\n                    cmap=\"gray\",\n                    vmin=0,\n                    vmax=1\n                )\n\n            if col_idx == 0:\n\n                ax.set_ylabel(\n                    f\"{plane}\\n{strategy_name}\",\n                    fontsize=11\n                )\n\n\nfig.suptitle(\n    \"Slice Sampling Comparison\",\n    fontsize=16\n)\n\nplt.tight_layout()\n\nOUTPUT_IMAGE = os.path.join(\n    AUDIT_DIR,\n    \"cell32_slice_sampling_comparison.png\"\n)\n\nplt.savefig(\n    OUTPUT_IMAGE,\n    dpi=150,\n    bbox_inches=\"tight\"\n)\n\nplt.show()\n\nplt.close(fig)\n\n\n# ---------------------------------------------------------------\n# Save strategy table\n# ---------------------------------------------------------------\n\nresults_df = pd.DataFrame(\n    results\n)\n\nOUTPUT_CSV = os.path.join(\n    AUDIT_DIR,\n    \"cell32_slice_sampling_indices.csv\"\n)\n\nresults_df.to_csv(\n    OUTPUT_CSV,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Final\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 32 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved visualization:\",\n    OUTPUT_IMAGE\n)\n\nprint(\n    \"Saved sampling table:\",\n    OUTPUT_CSV\n)\n\nprint()\nprint(\n    \"Compare Uniform vs Center vs Hybrid \"\n    \"in the generated image.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:52:27.14496Z","iopub.execute_input":"2026-08-11T08:52:27.146033Z","iopub.status.idle":"2026-08-11T08:52:37.639057Z","shell.execute_reply.started":"2026-08-11T08:52:27.145997Z","shell.execute_reply":"2026-08-11T08:52:37.637679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 33 - PYTORCH DATASET + DATALOADER\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODEL_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nFOLD_TABLE_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_study_folds.csv\"\n)\n\nIMAGE_SIZE = 224\n\nSLICES_PER_PLANE = 7\n\nPLANES = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\"\n]\n\nBATCH_SIZE = 2\n\nNUM_WORKERS = 0\n\nTARGET_COLUMNS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\n\nprint(\"=\" * 70)\nprint(\"CELL 33 - PYTORCH DATASET + DATALOADER\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Verify required functions from previous cells\n# ---------------------------------------------------------------\n\nrequired_functions = [\n    \"get_slice_position\",\n    \"normalize_dicom_image\",\n    \"resize_image\"\n]\n\nmissing_functions = [\n    name\n    for name in required_functions\n    if name not in globals()\n]\n\nif missing_functions:\n\n    raise RuntimeError(\n        \"Missing functions from previous cells: \"\n        + str(missing_functions)\n    )\n\n\n# ---------------------------------------------------------------\n# Hybrid sampling\n# ---------------------------------------------------------------\n\ndef hybrid_indices_7(\n    n,\n    k=SLICES_PER_PLANE,\n    margin_fraction=0.10\n):\n    \"\"\"\n    Select k slices while avoiding the extreme\n    first/last portions of the series.\n    \"\"\"\n\n    if n <= k:\n        return list(range(n))\n\n    margin = int(\n        round(\n            n * margin_fraction\n        )\n    )\n\n    start = margin\n    end = n - 1 - margin\n\n    if end <= start:\n\n        indices = np.linspace(\n            0,\n            n - 1,\n            k\n        )\n\n    else:\n\n        indices = np.linspace(\n            start,\n            end,\n            k\n        )\n\n    indices = np.round(\n        indices\n    ).astype(int)\n\n    return list(\n        np.unique(indices)\n    )\n\n\n# ---------------------------------------------------------------\n# DICOM file discovery\n# ---------------------------------------------------------------\n\ndef get_sorted_dicom_files_33(\n    study_id,\n    series_id\n):\n\n    series_path = os.path.join(\n        TRAIN_SERIES_DIR,\n        str(study_id),\n        str(series_id)\n    )\n\n    if not os.path.isdir(series_path):\n\n        raise FileNotFoundError(\n            f\"Series directory not found: {series_path}\"\n        )\n\n    files = []\n\n    for filename in os.listdir(\n        series_path\n    ):\n\n        path = os.path.join(\n            series_path,\n            filename\n        )\n\n        if os.path.isfile(path):\n\n            files.append(path)\n\n    records = []\n\n    for path in files:\n\n        try:\n\n            ds = pydicom.dcmread(\n                path,\n                stop_before_pixels=True,\n                force=True\n            )\n\n            position = get_slice_position(\n                ds\n            )\n\n            records.append(\n                (\n                    position,\n                    path\n                )\n            )\n\n        except Exception:\n\n            continue\n\n    records.sort(\n        key=lambda x: x[0]\n    )\n\n    return [\n        path\n        for _, path in records\n    ]\n\n\n# ---------------------------------------------------------------\n# Load selected DICOM slices\n# ---------------------------------------------------------------\n\ndef load_series_33(\n    study_id,\n    series_id,\n    num_slices=SLICES_PER_PLANE\n):\n\n    files = get_sorted_dicom_files_33(\n        study_id,\n        series_id\n    )\n\n    if len(files) == 0:\n\n        raise RuntimeError(\n            f\"No readable DICOM files: \"\n            f\"{study_id}/{series_id}\"\n        )\n\n    indices = hybrid_indices_7(\n        len(files),\n        num_slices\n    )\n\n    images = []\n\n    for index in indices:\n\n        ds = pydicom.dcmread(\n            files[index],\n            force=True\n        )\n\n        image = ds.pixel_array.astype(\n            np.float32\n        )\n\n        slope = float(\n            getattr(\n                ds,\n                \"RescaleSlope\",\n                1.0\n            )\n        )\n\n        intercept = float(\n            getattr(\n                ds,\n                \"RescaleIntercept\",\n                0.0\n            )\n        )\n\n        image = (\n            image * slope\n            + intercept\n        )\n\n        photometric = getattr(\n            ds,\n            \"PhotometricInterpretation\",\n            \"MONOCHROME2\"\n        )\n\n        if photometric == \"MONOCHROME1\":\n\n            image = (\n                image.max()\n                - image\n            )\n\n        image = normalize_dicom_image(\n            image\n        )\n\n        image = resize_image(\n            image,\n            IMAGE_SIZE\n        )\n\n        images.append(\n            image.astype(\n                np.float32\n            )\n        )\n\n    # If a very short series results in fewer\n    # unique indices, pad by repeating the\n    # last available image.\n\n    while len(images) < num_slices:\n\n        images.append(\n            images[-1].copy()\n        )\n\n    images = images[\n        :num_slices\n    ]\n\n    return np.stack(\n        images,\n        axis=0\n    ).astype(\n        np.float32\n    )\n\n\n# ---------------------------------------------------------------\n# Study-level dataset\n# ---------------------------------------------------------------\n\nclass KneeStudyDataset(\n    Dataset\n):\n\n    def __init__(\n        self,\n        dataframe,\n        targets=TARGET_COLUMNS\n    ):\n\n        self.df = dataframe.reset_index(\n            drop=True\n        )\n\n        self.targets = targets\n\n    def __len__(\n        self\n    ):\n\n        return len(self.df)\n\n    def __getitem__(\n        self,\n        index\n    ):\n\n        row = self.df.iloc[\n            index\n        ]\n\n        study_id = str(\n            row[\n                \"StudyInstanceUID\"\n            ]\n        )\n\n        plane_images = []\n\n        for plane in PLANES:\n\n            series_column = (\n                f\"{plane}_SeriesInstanceUID\"\n            )\n\n            series_id = str(\n                row[\n                    series_column\n                ]\n            )\n\n            image_stack = load_series_33(\n                study_id,\n                series_id,\n                SLICES_PER_PLANE\n            )\n\n            plane_images.append(\n                image_stack\n            )\n\n        # -------------------------------------------------------\n        # Shape:\n        # 3 planes × 7 slices × 224 × 224\n        # becomes:\n        # 21 × 224 × 224\n        # -------------------------------------------------------\n\n        image_tensor = np.concatenate(\n            plane_images,\n            axis=0\n        )\n\n        image_tensor = torch.from_numpy(\n            image_tensor\n        ).float()\n\n        labels = np.asarray(\n            [\n                float(row[target])\n                for target in self.targets\n            ],\n            dtype=np.float32\n        )\n\n        label_tensor = torch.from_numpy(\n            labels\n        ).float()\n\n        return {\n            \"image\": image_tensor,\n            \"label\": label_tensor,\n            \"study_id\": study_id\n        }\n\n\n# ---------------------------------------------------------------\n# Load modeling data\n# ---------------------------------------------------------------\n\nmodeling_df = pd.read_csv(\n    MODEL_TABLE_PATH\n)\n\nfold_df = pd.read_csv(\n    FOLD_TABLE_PATH\n)\n\n\nprint(\n    \"\\nModeling table:\",\n    modeling_df.shape\n)\n\nprint(\n    \"Fold table:\",\n    fold_df.shape\n)\n\n\n# ---------------------------------------------------------------\n# Merge fold information\n# ---------------------------------------------------------------\n\nif \"fold\" not in fold_df.columns:\n\n    raise ValueError(\n        \"Fold column not found in fold table.\"\n    )\n\n\nif (\n    \"StudyInstanceUID\"\n    not in fold_df.columns\n):\n\n    raise ValueError(\n        \"StudyInstanceUID missing from fold table.\"\n    )\n\n\nmodeling_df = modeling_df.drop(\n    columns=[\"fold\"],\n    errors=\"ignore\"\n)\n\nmodeling_df = modeling_df.merge(\n    fold_df[\n        [\n            \"StudyInstanceUID\",\n            \"fold\"\n        ]\n    ],\n    on=\"StudyInstanceUID\",\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\n\nif modeling_df[\"fold\"].isna().any():\n\n    raise RuntimeError(\n        \"Some modeling studies do not have a fold assignment.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Integrity checks\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"DATASET INTEGRITY\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Studies:\",\n    len(modeling_df)\n)\n\nprint(\n    \"Unique studies:\",\n    modeling_df[\n        \"StudyInstanceUID\"\n    ].nunique()\n)\n\nprint(\n    \"Fold distribution:\"\n)\n\nprint(\n    modeling_df[\n        \"fold\"\n    ].value_counts().sort_index()\n)\n\n\nfor plane in PLANES:\n\n    column = (\n        f\"{plane}_SeriesInstanceUID\"\n    )\n\n    missing = modeling_df[\n        column\n    ].isna().sum()\n\n    unique = modeling_df[\n        column\n    ].nunique()\n\n    print(\n        f\"{plane:8s}: \"\n        f\"missing={missing}, \"\n        f\"unique={unique}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Check targets\n# ---------------------------------------------------------------\n\nfor target in TARGET_COLUMNS:\n\n    values = set(\n        modeling_df[\n            target\n        ].dropna().unique()\n    )\n\n    if not values.issubset(\n        {0, 1}\n    ):\n\n        raise ValueError(\n            f\"Invalid values in target: {target}: {values}\"\n        )\n\n\nprint(\n    \"\\nTargets:\",\n    len(TARGET_COLUMNS)\n)\n\nprint(\n    \"All target columns contain only 0/1.\"\n)\n\n\n# ---------------------------------------------------------------\n# Create fold datasets\n# ---------------------------------------------------------------\n\ntrain_df = modeling_df[\n    modeling_df[\"fold\"] != 0\n].reset_index(\n    drop=True\n)\n\nval_df = modeling_df[\n    modeling_df[\"fold\"] == 0\n].reset_index(\n    drop=True\n)\n\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"FOLD 0 DATASET\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Training studies:\",\n    len(train_df)\n)\n\nprint(\n    \"Validation studies:\",\n    len(val_df)\n)\n\nprint(\n    \"Training/validation overlap:\",\n    len(\n        set(\n            train_df[\n                \"StudyInstanceUID\"\n            ]\n        )\n        &\n        set(\n            val_df[\n                \"StudyInstanceUID\"\n            ]\n        )\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Create datasets\n# ---------------------------------------------------------------\n\ntrain_dataset = KneeStudyDataset(\n    train_df\n)\n\nval_dataset = KneeStudyDataset(\n    val_df\n)\n\n\nprint()\nprint(\n    \"Train Dataset:\",\n    len(train_dataset)\n)\n\nprint(\n    \"Validation Dataset:\",\n    len(val_dataset)\n)\n\n\n# ---------------------------------------------------------------\n# Create DataLoaders\n# ---------------------------------------------------------------\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=False\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=False\n)\n\n\nprint()\nprint(\n    \"Train batches:\",\n    len(train_loader)\n)\n\nprint(\n    \"Validation batches:\",\n    len(val_loader)\n)\n\n\n# ---------------------------------------------------------------\n# Smoke test\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"DATALOADER SMOKE TEST\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\nbatch = next(\n    iter(train_loader)\n)\n\nimages = batch[\n    \"image\"\n]\n\nlabels = batch[\n    \"label\"\n]\n\nstudy_ids = batch[\n    \"study_id\"\n]\n\n\nprint(\n    \"Image batch shape:\",\n    tuple(images.shape)\n)\n\nprint(\n    \"Expected image shape:\",\n    (\n        BATCH_SIZE,\n        21,\n        IMAGE_SIZE,\n        IMAGE_SIZE\n    )\n)\n\nprint(\n    \"Label batch shape:\",\n    tuple(labels.shape)\n)\n\nprint(\n    \"Expected label shape:\",\n    (\n        BATCH_SIZE,\n        len(TARGET_COLUMNS)\n    )\n)\n\nprint(\n    \"Image dtype:\",\n    images.dtype\n)\n\nprint(\n    \"Label dtype:\",\n    labels.dtype\n)\n\nprint(\n    \"Image min:\",\n    float(images.min())\n)\n\nprint(\n    \"Image max:\",\n    float(images.max())\n)\n\nprint(\n    \"Image mean:\",\n    float(images.mean())\n)\n\nprint(\n    \"Image std:\",\n    float(images.std())\n)\n\nprint(\n    \"First study IDs:\",\n    list(study_ids)\n)\n\n\n# ---------------------------------------------------------------\n# Final shape validation\n# ---------------------------------------------------------------\n\nexpected_image_shape = (\n    BATCH_SIZE,\n    21,\n    IMAGE_SIZE,\n    IMAGE_SIZE\n)\n\nexpected_label_shape = (\n    BATCH_SIZE,\n    len(TARGET_COLUMNS)\n)\n\n\nif tuple(images.shape) != expected_image_shape:\n\n    raise RuntimeError(\n        f\"Unexpected image shape: \"\n        f\"{tuple(images.shape)}\"\n    )\n\n\nif tuple(labels.shape) != expected_label_shape:\n\n    raise RuntimeError(\n        f\"Unexpected label shape: \"\n        f\"{tuple(labels.shape)}\"\n    )\n\n\nif (\n    float(images.min()) < 0.0\n    or float(images.max()) > 1.0\n):\n\n    raise RuntimeError(\n        \"Images are outside expected [0, 1] range.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Memory estimate\n# ---------------------------------------------------------------\n\nbytes_per_study = (\n    21\n    * IMAGE_SIZE\n    * IMAGE_SIZE\n    * 4\n)\n\nmb_per_study = (\n    bytes_per_study\n    / (1024 ** 2)\n)\n\nbatch_mb = (\n    mb_per_study\n    * BATCH_SIZE\n)\n\nprint()\nprint(\n    \"Approximate image memory per study:\",\n    round(mb_per_study, 3),\n    \"MB\"\n)\n\nprint(\n    \"Approximate image memory per batch:\",\n    round(batch_mb, 3),\n    \"MB\"\n)\n\n\n# ---------------------------------------------------------------\n# Final\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"CELL 33 COMPLETE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Hybrid sampling:\",\n    SLICES_PER_PLANE,\n    \"slices per plane\"\n)\n\nprint(\n    \"Total model input:\",\n    \"21 x 224 x 224\"\n)\n\nprint(\n    \"Dataset loading is lazy.\"\n)\n\nprint(\n    \"Full dataset was NOT loaded into memory.\"\n)\n\nprint(\n    \"Fold-aware train/validation separation verified.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:55:23.219967Z","iopub.execute_input":"2026-08-11T08:55:23.22032Z","iopub.status.idle":"2026-08-11T08:55:24.955405Z","shell.execute_reply.started":"2026-08-11T08:55:23.220289Z","shell.execute_reply":"2026-08-11T08:55:24.954257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 34 - MODEL ARCHITECTURE + FORWARD PASS\n# ================================================================\n\nimport os\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision.models import (\n    resnet18,\n    ResNet18_Weights\n)\n\n\nprint(\"=\" * 70)\nprint(\"CELL 34 - MODEL ARCHITECTURE\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nNUM_INPUT_CHANNELS = 21\nNUM_CLASSES = len(TARGET_COLUMNS)\n\nIMAGE_SIZE = 224\n\nDEVICE = torch.device(\n    \"cuda\" if torch.cuda.is_available()\n    else \"cpu\"\n)\n\n\nprint()\nprint(\"PyTorch version:\", torch.__version__)\nprint(\"Torchvision version:\", torchvision.__version__)\nprint(\"Device:\", DEVICE)\nprint(\"Input channels:\", NUM_INPUT_CHANNELS)\nprint(\"Output classes:\", NUM_CLASSES)\n\nprint()\nprint(\"Target order:\")\n\nfor i, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    print(\n        f\"{i:2d}: {target}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Build pretrained ResNet18\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BUILDING RESNET18 BACKBONE\")\nprint(\"=\" * 70)\n\n\nweights = None\n\ntry:\n\n    weights = ResNet18_Weights.DEFAULT\n\n    backbone = resnet18(\n        weights=weights\n    )\n\n    pretrained_available = True\n\n    print(\n        \"Pretrained ResNet18 weights loaded.\"\n    )\n\nexcept Exception as e:\n\n    print(\n        \"WARNING: Pretrained weights could not be loaded.\"\n    )\n\n    print(\n        \"Reason:\",\n        str(e)\n    )\n\n    print(\n        \"Falling back to randomly initialized ResNet18.\"\n    )\n\n    backbone = resnet18(\n        weights=None\n    )\n\n    pretrained_available = False\n\n\n# ---------------------------------------------------------------\n# Modify first convolution\n#\n# Original:\n#     3 input channels\n#\n# New:\n#     21 input channels\n#\n# The original RGB filters are averaged and repeated\n# across the 21 medical-image channels.\n# ---------------------------------------------------------------\n\nold_conv = backbone.conv1\n\nnew_conv = nn.Conv2d(\n    in_channels=NUM_INPUT_CHANNELS,\n    out_channels=old_conv.out_channels,\n    kernel_size=old_conv.kernel_size,\n    stride=old_conv.stride,\n    padding=old_conv.padding,\n    bias=False\n)\n\n\nwith torch.no_grad():\n\n    if pretrained_available:\n\n        original_weight = (\n            old_conv.weight.data\n        )\n\n        mean_weight = (\n            original_weight.mean(\n                dim=1,\n                keepdim=True\n            )\n        )\n\n        repeated_weight = (\n            mean_weight.repeat(\n                1,\n                NUM_INPUT_CHANNELS,\n                1,\n                1\n            )\n        )\n\n        # Scale so the magnitude is comparable\n        # to the original convolution.\n\n        repeated_weight = (\n            repeated_weight\n            / NUM_INPUT_CHANNELS\n            * 3.0\n        )\n\n        new_conv.weight.copy_(\n            repeated_weight\n        )\n\n    else:\n\n        nn.init.kaiming_normal_(\n            new_conv.weight,\n            mode=\"fan_out\",\n            nonlinearity=\"relu\"\n        )\n\n\nbackbone.conv1 = new_conv\n\n\n# ---------------------------------------------------------------\n# Remove original ImageNet classifier\n# ---------------------------------------------------------------\n\nfeature_dim = (\n    backbone.fc.in_features\n)\n\nbackbone.fc = nn.Identity()\n\n\n# ---------------------------------------------------------------\n# Multilabel knee abnormality model\n# ---------------------------------------------------------------\n\nclass KneeAbnormalityModel(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        backbone,\n        feature_dim,\n        num_classes\n    ):\n\n        super().__init__()\n\n        self.backbone = backbone\n\n        self.dropout = nn.Dropout(\n            p=0.30\n        )\n\n        self.classifier = nn.Sequential(\n\n            nn.Linear(\n                feature_dim,\n                256\n            ),\n\n            nn.ReLU(\n                inplace=True\n            ),\n\n            nn.Dropout(\n                p=0.30\n            ),\n\n            nn.Linear(\n                256,\n                num_classes\n            )\n        )\n\n    def forward(\n        self,\n        x\n    ):\n\n        features = self.backbone(\n            x\n        )\n\n        features = self.dropout(\n            features\n        )\n\n        logits = self.classifier(\n            features\n        )\n\n        return logits\n\n\nmodel = KneeAbnormalityModel(\n    backbone=backbone,\n    feature_dim=feature_dim,\n    num_classes=NUM_CLASSES\n)\n\n\n# ---------------------------------------------------------------\n# Freeze backbone initially\n# ---------------------------------------------------------------\n\nfor parameter in model.backbone.parameters():\n\n    parameter.requires_grad = False\n\n\n# Keep the modified first convolution trainable.\n# This allows the network to adapt the ImageNet\n# input representation to the 21-channel MRI input.\n\nfor parameter in (\n    model.backbone.conv1.parameters()\n):\n\n    parameter.requires_grad = True\n\n\n# ---------------------------------------------------------------\n# Move model to device\n# ---------------------------------------------------------------\n\nmodel = model.to(\n    DEVICE\n)\n\n\n# ---------------------------------------------------------------\n# Parameter statistics\n# ---------------------------------------------------------------\n\ntotal_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n)\n\ntrainable_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\nfrozen_parameters = (\n    total_parameters\n    - trainable_parameters\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"MODEL PARAMETER SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    \"Total parameters:\",\n    f\"{total_parameters:,}\"\n)\n\nprint(\n    \"Trainable parameters:\",\n    f\"{trainable_parameters:,}\"\n)\n\nprint(\n    \"Frozen parameters:\",\n    f\"{frozen_parameters:,}\"\n)\n\nprint(\n    \"Trainable percentage:\",\n    f\"{100 * trainable_parameters / total_parameters:.2f}%\"\n)\n\n\n# ---------------------------------------------------------------\n# Forward-pass smoke test\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FORWARD PASS TEST\")\nprint(\"=\" * 70)\n\n\nmodel.eval()\n\n\nwith torch.no_grad():\n\n    sample_batch = next(\n        iter(train_loader)\n    )\n\n    sample_images = (\n        sample_batch[\"image\"]\n        .to(DEVICE)\n    )\n\n    sample_labels = (\n        sample_batch[\"label\"]\n        .to(DEVICE)\n    )\n\n    output_logits = model(\n        sample_images\n    )\n\n\nprint(\n    \"Input shape:\",\n    tuple(\n        sample_images.shape\n    )\n)\n\nprint(\n    \"Label shape:\",\n    tuple(\n        sample_labels.shape\n    )\n)\n\nprint(\n    \"Output logits shape:\",\n    tuple(\n        output_logits.shape\n    )\n)\n\nprint(\n    \"Expected output shape:\",\n    (\n        sample_images.shape[0],\n        NUM_CLASSES\n    )\n)\n\nprint(\n    \"Logit dtype:\",\n    output_logits.dtype\n)\n\nprint(\n    \"Logit minimum:\",\n    float(\n        output_logits.min()\n    )\n)\n\nprint(\n    \"Logit maximum:\",\n    float(\n        output_logits.max()\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Sigmoid probability test\n# ---------------------------------------------------------------\n\nprobabilities = torch.sigmoid(\n    output_logits\n)\n\nprint()\nprint(\n    \"Probability shape:\",\n    tuple(\n        probabilities.shape\n    )\n)\n\nprint(\n    \"Probability minimum:\",\n    float(\n        probabilities.min()\n    )\n)\n\nprint(\n    \"Probability maximum:\",\n    float(\n        probabilities.max()\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Output class mapping\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"OUTPUT MAPPING\")\nprint(\"=\" * 70)\n\nfor i, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    print(\n        f\"Output {i:2d} -> {target}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Final validation\n# ---------------------------------------------------------------\n\nexpected_output_shape = (\n    sample_images.shape[0],\n    NUM_CLASSES\n)\n\nif tuple(\n    output_logits.shape\n) != expected_output_shape:\n\n    raise RuntimeError(\n        \"Model output shape is incorrect.\"\n    )\n\n\nif not torch.isfinite(\n    output_logits\n).all():\n\n    raise RuntimeError(\n        \"Model produced NaN or infinite logits.\"\n    )\n\n\nif not torch.isfinite(\n    probabilities\n).all():\n\n    raise RuntimeError(\n        \"Model produced NaN or infinite probabilities.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Final\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 34 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Model construction: PASS\"\n)\n\nprint(\n    \"21-channel input: PASS\"\n)\n\nprint(\n    \"12-class multilabel output: PASS\"\n)\n\nprint(\n    \"Forward pass: PASS\"\n)\n\nprint(\n    \"NaN/Inf check: PASS\"\n)\n\nprint(\n    \"Backbone initially frozen: PASS\"\n)\n\nprint(\n    \"Ready for loss-function and training setup.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:56:41.78035Z","iopub.execute_input":"2026-08-11T08:56:41.780721Z","iopub.status.idle":"2026-08-11T08:56:48.678694Z","shell.execute_reply.started":"2026-08-11T08:56:41.78069Z","shell.execute_reply":"2026-08-11T08:56:48.677432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 35 - LOSS FUNCTION + CLASS IMBALANCE + OPTIMIZER\n# ================================================================\n\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\n\nprint(\"=\" * 70)\nprint(\"CELL 35 - LOSS + CLASS WEIGHTS + OPTIMIZER\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Reproducibility\n# ---------------------------------------------------------------\n\nSEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\n\nprint()\nprint(\"Random seed:\", SEED)\n\n\n# ---------------------------------------------------------------\n# Training configuration\n# ---------------------------------------------------------------\n\nLEARNING_RATE = 1e-3\n\nWEIGHT_DECAY = 1e-4\n\nMAX_GRAD_NORM = 1.0\n\nNUM_EPOCHS = 10\n\nEARLY_STOPPING_PATIENCE = 3\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"TRAINING CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(\"Learning rate:\", LEARNING_RATE)\nprint(\"Weight decay:\", WEIGHT_DECAY)\nprint(\"Maximum gradient norm:\", MAX_GRAD_NORM)\nprint(\"Maximum epochs:\", NUM_EPOCHS)\nprint(\n    \"Early stopping patience:\",\n    EARLY_STOPPING_PATIENCE\n)\n\n\n# ---------------------------------------------------------------\n# Calculate class imbalance weights\n#\n# IMPORTANT:\n# Calculate weights ONLY from the training split.\n# Validation labels must never influence training.\n# ---------------------------------------------------------------\n\ntrain_labels = train_df[\n    TARGET_COLUMNS\n].astype(\n    np.float32\n)\n\n\npositive_counts = (\n    train_labels.sum(\n        axis=0\n    )\n)\n\nnegative_counts = (\n    len(train_labels)\n    - positive_counts\n)\n\n\n# pos_weight = negative / positive\n#\n# BCEWithLogitsLoss then gives more importance to\n# positive examples for rare abnormalities.\n\npos_weight_values = (\n    negative_counts\n    / positive_counts\n)\n\n\npos_weight = torch.tensor(\n    pos_weight_values.values,\n    dtype=torch.float32,\n    device=DEVICE\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"CLASS IMBALANCE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"{'Target':22s}\"\n    f\"{'Positive':>10s}\"\n    f\"{'Negative':>10s}\"\n    f\"{'Pos Weight':>14s}\"\n)\n\nprint(\"-\" * 60)\n\nfor target, pos, neg, weight in zip(\n    TARGET_COLUMNS,\n    positive_counts,\n    negative_counts,\n    pos_weight_values\n):\n\n    print(\n        f\"{target:22s}\"\n        f\"{int(pos):10d}\"\n        f\"{int(neg):10d}\"\n        f\"{weight:14.4f}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Validate class weights\n# ---------------------------------------------------------------\n\nif not torch.isfinite(\n    pos_weight\n).all():\n\n    raise RuntimeError(\n        \"Invalid NaN/Inf values found in pos_weight.\"\n    )\n\n\nif (\n    pos_weight <= 0\n).any():\n\n    raise RuntimeError(\n        \"Invalid non-positive class weight found.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Loss function\n# ---------------------------------------------------------------\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"LOSS FUNCTION\")\nprint(\"=\" * 70)\n\nprint(\n    \"Loss:\",\n    criterion\n)\n\nprint(\n    \"Type:\",\n    type(criterion).__name__\n)\n\nprint(\n    \"Using per-target positive weights: YES\"\n)\n\n\n# ---------------------------------------------------------------\n# Optimizer\n#\n# Only currently trainable parameters are passed.\n# The frozen ResNet parameters will not be updated.\n# ---------------------------------------------------------------\n\ntrainable_parameters = [\n    parameter\n    for parameter in model.parameters()\n    if parameter.requires_grad\n]\n\n\noptimizer = torch.optim.AdamW(\n    trainable_parameters,\n    lr=LEARNING_RATE,\n    weight_decay=WEIGHT_DECAY\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"OPTIMIZER\")\nprint(\"=\" * 70)\n\nprint(\n    \"Optimizer:\",\n    type(optimizer).__name__\n)\n\nprint(\n    \"Trainable parameter tensors:\",\n    len(trainable_parameters)\n)\n\nprint(\n    \"Learning rate:\",\n    LEARNING_RATE\n)\n\nprint(\n    \"Weight decay:\",\n    WEIGHT_DECAY\n)\n\n\n# ---------------------------------------------------------------\n# Learning-rate scheduler\n#\n# Reduce LR when validation performance stops improving.\n# The scheduler will be stepped after validation.\n# ---------------------------------------------------------------\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"min\",\n    factor=0.5,\n    patience=1\n)\n\n\nprint(\n    \"Scheduler:\",\n    type(scheduler).__name__\n)\n\n\n# ---------------------------------------------------------------\n# Verify optimizer parameters\n# ---------------------------------------------------------------\n\noptimizer_parameter_count = sum(\n    parameter.numel()\n    for group in optimizer.param_groups\n    for parameter in group[\"params\"]\n)\n\n\nactual_trainable_count = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\n\nprint()\nprint(\n    \"Optimizer parameter count:\",\n    f\"{optimizer_parameter_count:,}\"\n)\n\nprint(\n    \"Actual trainable parameter count:\",\n    f\"{actual_trainable_count:,}\"\n)\n\n\nif (\n    optimizer_parameter_count\n    != actual_trainable_count\n):\n\n    raise RuntimeError(\n        \"Optimizer parameter count does not match \"\n        \"the model's trainable parameter count.\"\n    )\n\n\n# ---------------------------------------------------------------\n# One-batch loss test\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LOSS SMOKE TEST\")\nprint(\"=\" * 70)\n\n\nmodel.train()\n\n\nsample_batch = next(\n    iter(train_loader)\n)\n\nsample_images = (\n    sample_batch[\"image\"]\n    .to(DEVICE)\n)\n\nsample_labels = (\n    sample_batch[\"label\"]\n    .to(DEVICE)\n)\n\n\noptimizer.zero_grad(\n    set_to_none=True\n)\n\n\nsample_logits = model(\n    sample_images\n)\n\n\nsample_loss = criterion(\n    sample_logits,\n    sample_labels\n)\n\n\nprint(\n    \"Input shape:\",\n    tuple(\n        sample_images.shape\n    )\n)\n\nprint(\n    \"Logit shape:\",\n    tuple(\n        sample_logits.shape\n    )\n)\n\nprint(\n    \"Label shape:\",\n    tuple(\n        sample_labels.shape\n    )\n)\n\nprint(\n    \"Loss value:\",\n    float(\n        sample_loss\n    )\n)\n\n\nif not torch.isfinite(\n    sample_loss\n):\n\n    raise RuntimeError(\n        \"Loss produced NaN or Inf.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Backward-pass test\n# ---------------------------------------------------------------\n\nsample_loss.backward()\n\n\n# ---------------------------------------------------------------\n# Gradient statistics\n# ---------------------------------------------------------------\n\ngradient_norms = []\n\nfor parameter in model.parameters():\n\n    if (\n        parameter.requires_grad\n        and parameter.grad is not None\n    ):\n\n        gradient_norms.append(\n            parameter.grad.detach().norm().item()\n        )\n\n\nif len(gradient_norms) == 0:\n\n    raise RuntimeError(\n        \"No gradients were produced.\"\n    )\n\n\nmax_gradient = max(\n    gradient_norms\n)\n\nmean_gradient = np.mean(\n    gradient_norms\n)\n\n\nprint()\nprint(\n    \"Gradient tensors:\",\n    len(gradient_norms)\n)\n\nprint(\n    \"Mean gradient norm:\",\n    float(mean_gradient)\n)\n\nprint(\n    \"Maximum gradient norm:\",\n    float(max_gradient)\n)\n\n\nif not np.isfinite(\n    max_gradient\n):\n\n    raise RuntimeError(\n        \"Invalid gradient detected.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Clear gradients after smoke test\n# ---------------------------------------------------------------\n\noptimizer.zero_grad(\n    set_to_none=True\n)\n\n\n# ---------------------------------------------------------------\n# Final configuration summary\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 35 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Weighted BCE loss: PASS\"\n)\n\nprint(\n    \"Class imbalance handling: PASS\"\n)\n\nprint(\n    \"Optimizer: PASS\"\n)\n\nprint(\n    \"Scheduler: PASS\"\n)\n\nprint(\n    \"Backward pass: PASS\"\n)\n\nprint(\n    \"Gradient validity: PASS\"\n)\n\nprint()\nprint(\n    \"Ready for one-batch training test.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T08:58:43.817351Z","iopub.execute_input":"2026-08-11T08:58:43.818393Z","iopub.status.idle":"2026-08-11T08:58:47.683366Z","shell.execute_reply.started":"2026-08-11T08:58:43.81835Z","shell.execute_reply":"2026-08-11T08:58:47.682481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 36 - ONE-BATCH TRAINING + CPU PERFORMANCE TEST\n# ================================================================\n\nimport time\nimport numpy as np\nimport torch\n\n\nprint(\"=\" * 70)\nprint(\"CELL 36 - ONE-BATCH TRAINING TEST\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nGRADIENT_CLIP_NORM = MAX_GRAD_NORM\n\n\nprint()\nprint(\"Device:\", DEVICE)\nprint(\"Gradient clipping:\", GRADIENT_CLIP_NORM)\n\n\n# ---------------------------------------------------------------\n# Prepare model\n# ---------------------------------------------------------------\n\nmodel.train()\n\noptimizer.zero_grad(\n    set_to_none=True\n)\n\n\n# ---------------------------------------------------------------\n# Load one batch\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LOADING TRAINING BATCH\")\nprint(\"=\" * 70)\n\n\nbatch_start = time.perf_counter()\n\n\nbatch = next(\n    iter(train_loader)\n)\n\n\nbatch_load_time = (\n    time.perf_counter()\n    - batch_start\n)\n\n\nimages = batch[\n    \"image\"\n].to(\n    DEVICE\n)\n\nlabels = batch[\n    \"label\"\n].to(\n    DEVICE\n)\n\n\nprint(\n    \"Batch loaded in:\",\n    f\"{batch_load_time:.3f} seconds\"\n)\n\nprint(\n    \"Images:\",\n    tuple(images.shape)\n)\n\nprint(\n    \"Labels:\",\n    tuple(labels.shape)\n)\n\n\n# ---------------------------------------------------------------\n# Forward pass\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"FORWARD + BACKWARD + OPTIMIZER STEP\"\n)\n\nprint(\n    \"=\" * 70\n)\n\n\nstart_time = time.perf_counter()\n\n\nlogits = model(\n    images\n)\n\n\nforward_time = (\n    time.perf_counter()\n    - start_time\n)\n\n\n# ---------------------------------------------------------------\n# Loss\n# ---------------------------------------------------------------\n\nloss = criterion(\n    logits,\n    labels\n)\n\n\nif not torch.isfinite(\n    loss\n):\n\n    raise RuntimeError(\n        \"Training loss is NaN or Inf.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Backward\n# ---------------------------------------------------------------\n\nbackward_start = time.perf_counter()\n\n\nloss.backward()\n\n\nbackward_time = (\n    time.perf_counter()\n    - backward_start\n)\n\n\n# ---------------------------------------------------------------\n# Gradient statistics BEFORE clipping\n# ---------------------------------------------------------------\n\ngradient_norm_before = (\n    torch.nn.utils.clip_grad_norm_(\n        model.parameters(),\n        max_norm=GRADIENT_CLIP_NORM\n    )\n)\n\n\ngradient_norm_before_value = float(\n    gradient_norm_before\n)\n\n\nprint(\n    \"Gradient norm before clipping:\",\n    gradient_norm_before_value\n)\n\n\n# ---------------------------------------------------------------\n# Optimizer step\n# ---------------------------------------------------------------\n\noptimizer_start = time.perf_counter()\n\n\noptimizer.step()\n\n\noptimizer_time = (\n    time.perf_counter()\n    - optimizer_start\n)\n\n\ntotal_training_time = (\n    time.perf_counter()\n    - start_time\n)\n\n\n# ---------------------------------------------------------------\n# Clear gradients\n# ---------------------------------------------------------------\n\noptimizer.zero_grad(\n    set_to_none=True\n)\n\n\n# ---------------------------------------------------------------\n# Check parameters remain valid\n# ---------------------------------------------------------------\n\ninvalid_parameters = 0\n\nfor parameter in model.parameters():\n\n    if not torch.isfinite(\n        parameter\n    ).all():\n\n        invalid_parameters += 1\n\n\nif invalid_parameters > 0:\n\n    raise RuntimeError(\n        \"Invalid NaN/Inf model parameters detected \"\n        f\"after optimizer step: {invalid_parameters}\"\n    )\n\n\n# ---------------------------------------------------------------\n# Probability check\n# ---------------------------------------------------------------\n\nwith torch.no_grad():\n\n    probabilities = torch.sigmoid(\n        logits\n    )\n\n\nif not torch.isfinite(\n    probabilities\n).all():\n\n    raise RuntimeError(\n        \"Invalid probabilities detected.\"\n    )\n\n\n# ---------------------------------------------------------------\n# Timing summary\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"PERFORMANCE SUMMARY\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Batch loading time:\",\n    f\"{batch_load_time:.3f} sec\"\n)\n\nprint(\n    \"Forward time:\",\n    f\"{forward_time:.3f} sec\"\n)\n\nprint(\n    \"Backward time:\",\n    f\"{backward_time:.3f} sec\"\n)\n\nprint(\n    \"Optimizer step:\",\n    f\"{optimizer_time:.3f} sec\"\n)\n\nprint(\n    \"Total training step:\",\n    f\"{total_training_time:.3f} sec\"\n)\n\n\n# ---------------------------------------------------------------\n# Training result\n# ---------------------------------------------------------------\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"TRAINING STEP RESULT\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Loss:\",\n    float(loss)\n)\n\nprint(\n    \"Logit min:\",\n    float(logits.min())\n)\n\nprint(\n    \"Logit max:\",\n    float(logits.max())\n)\n\nprint(\n    \"Probability min:\",\n    float(probabilities.min())\n)\n\nprint(\n    \"Probability max:\",\n    float(probabilities.max())\n)\n\nprint(\n    \"Gradient clipping applied:\",\n    \"YES\"\n)\n\nprint(\n    \"Model parameters valid:\",\n    \"YES\"\n)\n\n\n# ---------------------------------------------------------------\n# Estimate Fold-0 epoch time\n# ---------------------------------------------------------------\n\nsteps_per_epoch = len(\n    train_loader\n)\n\nestimated_epoch_seconds = (\n    total_training_time\n    * steps_per_epoch\n)\n\n\nestimated_epoch_minutes = (\n    estimated_epoch_seconds\n    / 60.0\n)\n\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"TRAINING TIME ESTIMATE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Training batches per epoch:\",\n    steps_per_epoch\n)\n\nprint(\n    \"Estimated Fold-0 epoch time:\",\n    f\"{estimated_epoch_minutes:.2f} minutes\"\n)\n\nprint(\n    \"Estimated 10-epoch Fold-0 time:\",\n    f\"{estimated_epoch_minutes * 10:.2f} minutes\"\n)\n\n\n# ---------------------------------------------------------------\n# Final validation\n# ---------------------------------------------------------------\n\nif not np.isfinite(\n    float(loss)\n):\n\n    raise RuntimeError(\n        \"Invalid loss.\"\n    )\n\n\nif not np.isfinite(\n    gradient_norm_before_value\n):\n\n    raise RuntimeError(\n        \"Invalid gradient norm.\"\n    )\n\n\nif total_training_time <= 0:\n\n    raise RuntimeError(\n        \"Invalid timing measurement.\"\n    )\n\n\nprint()\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"CELL 36 COMPLETE\"\n)\n\nprint(\n    \"=\" * 70\n)\n\nprint(\n    \"Forward pass: PASS\"\n)\n\nprint(\n    \"Backward pass: PASS\"\n)\n\nprint(\n    \"Gradient clipping: PASS\"\n)\n\nprint(\n    \"Optimizer update: PASS\"\n)\n\nprint(\n    \"Parameter validity: PASS\"\n)\n\nprint(\n    \"CPU performance measurement: PASS\"\n)\n\nprint()\nprint(\n    \"Ready to decide the full training configuration.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:00:20.370003Z","iopub.execute_input":"2026-08-11T09:00:20.370344Z","iopub.status.idle":"2026-08-11T09:00:22.83923Z","shell.execute_reply.started":"2026-08-11T09:00:20.370313Z","shell.execute_reply":"2026-08-11T09:00:22.838115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 37 - FOLD 0 TRAINING\n# ================================================================\n\nimport os\nimport time\nimport copy\nimport numpy as np\nimport pandas as pd\nimport torch\n\n\nprint(\"=\" * 70)\nprint(\"CELL 37 - FOLD 0 TRAINING\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Configuration\n# ---------------------------------------------------------------\n\nFOLD_ID = 0\n\nNUM_EPOCHS = 10\n\nEARLY_STOPPING_PATIENCE = 3\n\nGRADIENT_CLIP_NORM = 1.0\n\nCHECKPOINT_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell37_fold0_best_model.pt\"\n)\n\nHISTORY_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell37_fold0_history.csv\"\n)\n\n\nprint()\nprint(\"Fold:\", FOLD_ID)\nprint(\"Maximum epochs:\", NUM_EPOCHS)\nprint(\n    \"Early stopping patience:\",\n    EARLY_STOPPING_PATIENCE\n)\n\nprint(\n    \"Device:\",\n    DEVICE\n)\n\n\n# ---------------------------------------------------------------\n# Recreate Fold-0 train/validation datasets\n# ---------------------------------------------------------------\n\nfold_train_df = modeling_df[\n    modeling_df[\"fold\"] != FOLD_ID\n].reset_index(\n    drop=True\n)\n\nfold_val_df = modeling_df[\n    modeling_df[\"fold\"] == FOLD_ID\n].reset_index(\n    drop=True\n)\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"FOLD DATA\")\nprint(\"=\" * 70)\n\nprint(\n    \"Training studies:\",\n    len(fold_train_df)\n)\n\nprint(\n    \"Validation studies:\",\n    len(fold_val_df)\n)\n\n\ntrain_dataset_fold0 = KneeStudyDataset(\n    fold_train_df\n)\n\nval_dataset_fold0 = KneeStudyDataset(\n    fold_val_df\n)\n\n\ntrain_loader_fold0 = DataLoader(\n    train_dataset_fold0,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=False\n)\n\nval_loader_fold0 = DataLoader(\n    val_dataset_fold0,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=False\n)\n\n\nprint(\n    \"Training batches:\",\n    len(train_loader_fold0)\n)\n\nprint(\n    \"Validation batches:\",\n    len(val_loader_fold0)\n)\n\n\n# ---------------------------------------------------------------\n# Reset model to the original Cell 34 state\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"RESETTING MODEL\")\nprint(\"=\" * 70)\n\n\nweights = ResNet18_Weights.DEFAULT\n\nbackbone = resnet18(\n    weights=weights\n)\n\n\nold_conv = backbone.conv1\n\nnew_conv = nn.Conv2d(\n    in_channels=NUM_INPUT_CHANNELS,\n    out_channels=old_conv.out_channels,\n    kernel_size=old_conv.kernel_size,\n    stride=old_conv.stride,\n    padding=old_conv.padding,\n    bias=False\n)\n\n\nwith torch.no_grad():\n\n    original_weight = (\n        old_conv.weight.data\n    )\n\n    mean_weight = (\n        original_weight.mean(\n            dim=1,\n            keepdim=True\n        )\n    )\n\n    repeated_weight = (\n        mean_weight.repeat(\n            1,\n            NUM_INPUT_CHANNELS,\n            1,\n            1\n        )\n    )\n\n    repeated_weight = (\n        repeated_weight\n        / NUM_INPUT_CHANNELS\n        * 3.0\n    )\n\n    new_conv.weight.copy_(\n        repeated_weight\n    )\n\n\nbackbone.conv1 = new_conv\n\nfeature_dim = (\n    backbone.fc.in_features\n)\n\nbackbone.fc = nn.Identity()\n\n\nmodel = KneeAbnormalityModel(\n    backbone=backbone,\n    feature_dim=feature_dim,\n    num_classes=NUM_CLASSES\n)\n\n\n# Freeze backbone\n\nfor parameter in model.backbone.parameters():\n\n    parameter.requires_grad = False\n\n\n# Keep first convolution trainable\n\nfor parameter in (\n    model.backbone.conv1.parameters()\n):\n\n    parameter.requires_grad = True\n\n\nmodel = model.to(\n    DEVICE\n)\n\n\n# ---------------------------------------------------------------\n# Fold-specific class weights\n# ---------------------------------------------------------------\n\nfold_train_labels = (\n    fold_train_df[\n        TARGET_COLUMNS\n    ].astype(\n        np.float32\n    )\n)\n\n\npositive_counts_fold = (\n    fold_train_labels.sum(\n        axis=0\n    )\n)\n\nnegative_counts_fold = (\n    len(fold_train_labels)\n    - positive_counts_fold\n)\n\n\npos_weight_fold = torch.tensor(\n    (\n        negative_counts_fold\n        / positive_counts_fold\n    ).values,\n    dtype=torch.float32,\n    device=DEVICE\n)\n\n\ncriterion_fold = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_fold\n)\n\n\n# ---------------------------------------------------------------\n# Optimizer\n# ---------------------------------------------------------------\n\noptimizer_fold = torch.optim.AdamW(\n    [\n        parameter\n        for parameter in model.parameters()\n        if parameter.requires_grad\n    ],\n    lr=LEARNING_RATE,\n    weight_decay=WEIGHT_DECAY\n)\n\n\nscheduler_fold = (\n    torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer_fold,\n        mode=\"min\",\n        factor=0.5,\n        patience=1\n    )\n)\n\n\nprint(\n    \"Trainable parameters:\",\n    f\"{sum(p.numel() for p in model.parameters() if p.requires_grad):,}\"\n)\n\n\n# ---------------------------------------------------------------\n# Validation function\n# ---------------------------------------------------------------\n\ndef evaluate_fold0(\n    model,\n    loader,\n    criterion,\n    device\n):\n\n    model.eval()\n\n    total_loss = 0.0\n\n    total_samples = 0\n\n    all_logits = []\n\n    all_labels = []\n\n    with torch.no_grad():\n\n        for batch in loader:\n\n            images = batch[\n                \"image\"\n            ].to(\n                device\n            )\n\n            labels = batch[\n                \"label\"\n            ].to(\n                device\n            )\n\n            logits = model(\n                images\n            )\n\n            loss = criterion(\n                logits,\n                labels\n            )\n\n            batch_size = (\n                images.shape[0]\n            )\n\n            total_loss += (\n                float(loss)\n                * batch_size\n            )\n\n            total_samples += (\n                batch_size\n            )\n\n            all_logits.append(\n                logits.cpu()\n            )\n\n            all_labels.append(\n                labels.cpu()\n            )\n\n\n    mean_loss = (\n        total_loss\n        / total_samples\n    )\n\n\n    logits_array = torch.cat(\n        all_logits,\n        dim=0\n    ).numpy()\n\n\n    labels_array = torch.cat(\n        all_labels,\n        dim=0\n    ).numpy()\n\n\n    probabilities = (\n        1.0\n        / (\n            1.0\n            + np.exp(\n                -np.clip(\n                    logits_array,\n                    -50,\n                    50\n                )\n            )\n        )\n    )\n\n\n    predictions = (\n        probabilities >= 0.5\n    ).astype(\n        np.int32\n    )\n\n\n    return (\n        mean_loss,\n        labels_array,\n        probabilities,\n        predictions\n    )\n\n\n# ---------------------------------------------------------------\n# Metric calculation\n# ---------------------------------------------------------------\n\ndef calculate_metrics(\n    labels,\n    probabilities,\n    predictions\n):\n\n    results = []\n\n    for index, target in enumerate(\n        TARGET_COLUMNS\n    ):\n\n        y_true = labels[\n            :,\n            index\n        ]\n\n        y_prob = probabilities[\n            :,\n            index\n        ]\n\n        y_pred = predictions[\n            :,\n            index\n        ]\n\n\n        tp = np.sum(\n            (\n                y_true == 1\n            )\n            &\n            (\n                y_pred == 1\n            )\n        )\n\n        tn = np.sum(\n            (\n                y_true == 0\n            )\n            &\n            (\n                y_pred == 0\n            )\n        )\n\n        fp = np.sum(\n            (\n                y_true == 0\n            )\n            &\n            (\n                y_pred == 1\n            )\n        )\n\n        fn = np.sum(\n            (\n                y_true == 1\n            )\n            &\n            (\n                y_pred == 0\n            )\n        )\n\n\n        sensitivity = (\n            tp / (tp + fn)\n            if (tp + fn) > 0\n            else np.nan\n        )\n\n        specificity = (\n            tn / (tn + fp)\n            if (tn + fp) > 0\n            else np.nan\n        )\n\n        precision = (\n            tp / (tp + fp)\n            if (tp + fp) > 0\n            else 0.0\n        )\n\n        f1 = (\n            2 * precision * sensitivity\n            / (precision + sensitivity)\n            if (\n                precision + sensitivity\n            ) > 0\n            else 0.0\n        )\n\n\n        results.append(\n            {\n                \"target\": target,\n                \"positive_count\": int(\n                    np.sum(\n                        y_true == 1\n                    )\n                ),\n                \"negative_count\": int(\n                    np.sum(\n                        y_true == 0\n                    )\n                ),\n                \"predicted_positive\": int(\n                    np.sum(\n                        y_pred == 1\n                    )\n                ),\n                \"sensitivity\": sensitivity,\n                \"specificity\": specificity,\n                \"precision\": precision,\n                \"f1\": f1\n            }\n        )\n\n\n    return pd.DataFrame(\n        results\n    )\n\n\n# ---------------------------------------------------------------\n# Training loop\n# ---------------------------------------------------------------\n\nhistory = []\n\nbest_val_loss = float(\n    \"inf\"\n)\n\nbest_epoch = None\n\nbest_state = None\n\nepochs_without_improvement = 0\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"STARTING FOLD 0 TRAINING\")\nprint(\"=\" * 70)\n\n\nfor epoch in range(\n    1,\n    NUM_EPOCHS + 1\n):\n\n    epoch_start = time.perf_counter()\n\n\n    # -----------------------------------------------------------\n    # TRAIN\n    # -----------------------------------------------------------\n\n    model.train()\n\n    train_loss_sum = 0.0\n\n    train_samples = 0\n\n    train_batch_times = []\n\n\n    for batch in train_loader_fold0:\n\n        batch_start = time.perf_counter()\n\n\n        images = batch[\n            \"image\"\n        ].to(\n            DEVICE\n        )\n\n        labels = batch[\n            \"label\"\n        ].to(\n            DEVICE\n        )\n\n\n        optimizer_fold.zero_grad(\n            set_to_none=True\n        )\n\n\n        logits = model(\n            images\n        )\n\n\n        loss = criterion_fold(\n            logits,\n            labels\n        )\n\n\n        if not torch.isfinite(\n            loss\n        ):\n\n            raise RuntimeError(\n                f\"Non-finite training loss \"\n                f\"at epoch {epoch}.\"\n            )\n\n\n        loss.backward()\n\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=GRADIENT_CLIP_NORM\n        )\n\n\n        optimizer_fold.step()\n\n\n        batch_size = (\n            images.shape[0]\n        )\n\n\n        train_loss_sum += (\n            float(loss)\n            * batch_size\n        )\n\n        train_samples += (\n            batch_size\n        )\n\n\n        train_batch_times.append(\n            time.perf_counter()\n            - batch_start\n        )\n\n\n    train_loss = (\n        train_loss_sum\n        / train_samples\n    )\n\n\n    # -----------------------------------------------------------\n    # VALIDATION\n    # -----------------------------------------------------------\n\n    validation_start = (\n        time.perf_counter()\n    )\n\n\n    (\n        val_loss,\n        val_labels,\n        val_probabilities,\n        val_predictions\n    ) = evaluate_fold0(\n        model,\n        val_loader_fold0,\n        criterion_fold,\n        DEVICE\n    )\n\n\n    validation_time = (\n        time.perf_counter()\n        - validation_start\n    )\n\n\n    # -----------------------------------------------------------\n    # Metrics\n    # -----------------------------------------------------------\n\n    metric_df = calculate_metrics(\n        val_labels,\n        val_probabilities,\n        val_predictions\n    )\n\n\n    mean_f1 = float(\n        metric_df[\n            \"f1\"\n        ].mean()\n    )\n\n\n    mean_sensitivity = float(\n        metric_df[\n            \"sensitivity\"\n        ].replace(\n            [np.inf, -np.inf],\n            np.nan\n        ).mean()\n    )\n\n\n    mean_specificity = float(\n        metric_df[\n            \"specificity\"\n        ].replace(\n            [np.inf, -np.inf],\n            np.nan\n        ).mean()\n    )\n\n\n    mean_precision = float(\n        metric_df[\n            \"precision\"\n        ].mean()\n    )\n\n\n    # -----------------------------------------------------------\n    # Scheduler\n    # -----------------------------------------------------------\n\n    scheduler_fold.step(\n        val_loss\n    )\n\n\n    current_lr = (\n        optimizer_fold\n        .param_groups[0][\"lr\"]\n    )\n\n\n    # -----------------------------------------------------------\n    # Epoch timing\n    # -----------------------------------------------------------\n\n    epoch_time = (\n        time.perf_counter()\n        - epoch_start\n    )\n\n\n    average_batch_time = (\n        np.mean(\n            train_batch_times\n        )\n    )\n\n\n    # -----------------------------------------------------------\n    # Improvement check\n    # -----------------------------------------------------------\n\n    improved = (\n        val_loss\n        < best_val_loss\n        - 1e-5\n    )\n\n\n    if improved:\n\n        best_val_loss = val_loss\n\n        best_epoch = epoch\n\n        best_state = copy.deepcopy(\n            model.state_dict()\n        )\n\n        torch.save(\n            {\n                \"epoch\": epoch,\n                \"model_state_dict\": best_state,\n                \"optimizer_state_dict\":\n                    optimizer_fold.state_dict(),\n                \"scheduler_state_dict\":\n                    scheduler_fold.state_dict(),\n                \"best_val_loss\":\n                    best_val_loss,\n                \"target_columns\":\n                    TARGET_COLUMNS,\n                \"fold\":\n                    FOLD_ID\n            },\n            CHECKPOINT_PATH\n        )\n\n        epochs_without_improvement = 0\n\n    else:\n\n        epochs_without_improvement += 1\n\n\n    # -----------------------------------------------------------\n    # Store history\n    # -----------------------------------------------------------\n\n    history.append(\n        {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"mean_f1\": mean_f1,\n            \"mean_sensitivity\":\n                mean_sensitivity,\n            \"mean_specificity\":\n                mean_specificity,\n            \"mean_precision\":\n                mean_precision,\n            \"learning_rate\": current_lr,\n            \"epoch_time_seconds\":\n                epoch_time,\n            \"validation_time_seconds\":\n                validation_time,\n            \"average_batch_time_seconds\":\n                average_batch_time\n        }\n    )\n\n\n    # -----------------------------------------------------------\n    # Print epoch summary\n    # -----------------------------------------------------------\n\n    print()\n    print(\n        f\"Epoch {epoch:02d}/{NUM_EPOCHS}\"\n    )\n\n    print(\n        f\"Train Loss: {train_loss:.5f}\"\n    )\n\n    print(\n        f\"Val Loss:   {val_loss:.5f}\"\n    )\n\n    print(\n        f\"Mean F1:    {mean_f1:.4f}\"\n    )\n\n    print(\n        f\"Sensitivity: {mean_sensitivity:.4f}\"\n    )\n\n    print(\n        f\"Specificity: {mean_specificity:.4f}\"\n    )\n\n    print(\n        f\"Precision:   {mean_precision:.4f}\"\n    )\n\n    print(\n        f\"Learning Rate: {current_lr:.6f}\"\n    )\n\n    print(\n        f\"Epoch Time: {epoch_time:.2f} sec\"\n    )\n\n    print(\n        f\"Avg Train Batch: \"\n        f\"{average_batch_time:.2f} sec\"\n    )\n\n    print(\n        f\"Validation Time: \"\n        f\"{validation_time:.2f} sec\"\n    )\n\n    print(\n        \"Best Epoch:\",\n        best_epoch\n    )\n\n    print(\n        \"Best Val Loss:\",\n        f\"{best_val_loss:.5f}\"\n    )\n\n\n    # -----------------------------------------------------------\n    # Early stopping\n    # -----------------------------------------------------------\n\n    if (\n        epochs_without_improvement\n        >= EARLY_STOPPING_PATIENCE\n    ):\n\n        print()\n        print(\n            \"Early stopping triggered.\"\n        )\n\n        break\n\n\n# ---------------------------------------------------------------\n# Restore best model\n# ---------------------------------------------------------------\n\nif best_state is not None:\n\n    model.load_state_dict(\n        best_state\n    )\n\n\n# ---------------------------------------------------------------\n# Save history\n# ---------------------------------------------------------------\n\nhistory_df = pd.DataFrame(\n    history\n)\n\nhistory_df.to_csv(\n    HISTORY_PATH,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Final validation using best model\n# ---------------------------------------------------------------\n\n(\n    final_val_loss,\n    final_val_labels,\n    final_val_probabilities,\n    final_val_predictions\n) = evaluate_fold0(\n    model,\n    val_loader_fold0,\n    criterion_fold,\n    DEVICE\n)\n\n\nfinal_metrics = calculate_metrics(\n    final_val_labels,\n    final_val_probabilities,\n    final_val_predictions\n)\n\n\n# ---------------------------------------------------------------\n# Save validation metrics\n# ---------------------------------------------------------------\n\nMETRICS_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell37_fold0_metrics.csv\"\n)\n\nfinal_metrics.to_csv(\n    METRICS_PATH,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Final summary\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FOLD 0 TRAINING COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Best epoch:\",\n    best_epoch\n)\n\nprint(\n    \"Best validation loss:\",\n    f\"{best_val_loss:.6f}\"\n)\n\nprint(\n    \"Final validation loss:\",\n    f\"{final_val_loss:.6f}\"\n)\n\nprint()\nprint(\"Per-target validation metrics:\")\n\nprint(\n    final_metrics.to_string(\n        index=False\n    )\n)\n\nprint()\nprint(\n    \"Mean F1:\",\n    f\"{final_metrics['f1'].mean():.4f}\"\n)\n\nprint(\n    \"Mean sensitivity:\",\n    f\"{final_metrics['sensitivity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean specificity:\",\n    f\"{final_metrics['specificity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean precision:\",\n    f\"{final_metrics['precision'].mean():.4f}\"\n)\n\nprint()\nprint(\n    \"Checkpoint saved:\",\n    CHECKPOINT_PATH\n)\n\nprint(\n    \"History saved:\",\n    HISTORY_PATH\n)\n\nprint(\n    \"Metrics saved:\",\n    METRICS_PATH\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 37 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:03:19.08149Z","iopub.execute_input":"2026-08-11T09:03:19.081846Z","iopub.status.idle":"2026-08-11T09:09:05.029324Z","shell.execute_reply.started":"2026-08-11T09:03:19.081814Z","shell.execute_reply":"2026-08-11T09:09:05.028171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 38 - FOLD 0 BEST MODEL DIAGNOSTIC\n# ================================================================\n\nprint(\"=\" * 70)\nprint(\"CELL 38 - FOLD 0 BEST MODEL DIAGNOSTIC\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load best checkpoint\n# ---------------------------------------------------------------\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\n\nprint()\nprint(\"Best checkpoint epoch:\")\nprint(\n    checkpoint[\"epoch\"]\n)\n\nprint(\n    \"Best validation loss:\",\n    checkpoint[\"best_val_loss\"]\n)\n\n\n# ---------------------------------------------------------------\n# Run validation\n# ---------------------------------------------------------------\n\n(\n    diagnostic_loss,\n    diagnostic_labels,\n    diagnostic_probabilities,\n    diagnostic_predictions\n) = evaluate_fold0(\n    model,\n    val_loader_fold0,\n    criterion_fold,\n    DEVICE\n)\n\n\n# ---------------------------------------------------------------\n# Build detailed diagnostics\n# ---------------------------------------------------------------\n\ndiagnostic_rows = []\n\n\nfor index, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    y_true = diagnostic_labels[\n        :,\n        index\n    ]\n\n    y_prob = diagnostic_probabilities[\n        :,\n        index\n    ]\n\n    y_pred = diagnostic_predictions[\n        :,\n        index\n    ]\n\n\n    actual_positive = int(\n        np.sum(\n            y_true == 1\n        )\n    )\n\n    actual_negative = int(\n        np.sum(\n            y_true == 0\n        )\n    )\n\n    predicted_positive = int(\n        np.sum(\n            y_pred == 1\n        )\n    )\n\n    predicted_negative = int(\n        np.sum(\n            y_pred == 0\n        )\n    )\n\n\n    tp = int(\n        np.sum(\n            (y_true == 1)\n            &\n            (y_pred == 1)\n        )\n    )\n\n    tn = int(\n        np.sum(\n            (y_true == 0)\n            &\n            (y_pred == 0)\n        )\n    )\n\n    fp = int(\n        np.sum(\n            (y_true == 0)\n            &\n            (y_pred == 1)\n        )\n    )\n\n    fn = int(\n        np.sum(\n            (y_true == 1)\n            &\n            (y_pred == 0)\n        )\n    )\n\n\n    sensitivity = (\n        tp / (tp + fn)\n        if (tp + fn) > 0\n        else np.nan\n    )\n\n    specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else np.nan\n    )\n\n    precision = (\n        tp / (tp + fp)\n        if (tp + fp) > 0\n        else 0.0\n    )\n\n    f1 = (\n        2 * precision * sensitivity\n        / (precision + sensitivity)\n        if (\n            precision + sensitivity\n        ) > 0\n        else 0.0\n    )\n\n\n    diagnostic_rows.append(\n        {\n            \"target\": target,\n            \"actual_positive\": actual_positive,\n            \"actual_negative\": actual_negative,\n            \"predicted_positive\": predicted_positive,\n            \"predicted_negative\": predicted_negative,\n            \"TP\": tp,\n            \"TN\": tn,\n            \"FP\": fp,\n            \"FN\": fn,\n            \"mean_probability\": float(\n                np.mean(y_prob)\n            ),\n            \"min_probability\": float(\n                np.min(y_prob)\n            ),\n            \"max_probability\": float(\n                np.max(y_prob)\n            ),\n            \"sensitivity\": sensitivity,\n            \"specificity\": specificity,\n            \"precision\": precision,\n            \"f1\": f1\n        }\n    )\n\n\ndiagnostic_df = pd.DataFrame(\n    diagnostic_rows\n)\n\n\n# ---------------------------------------------------------------\n# Save diagnostic table\n# ---------------------------------------------------------------\n\nDIAGNOSTIC_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell38_fold0_best_model_diagnostics.csv\"\n)\n\ndiagnostic_df.to_csv(\n    DIAGNOSTIC_PATH,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Print diagnostics\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BEST MODEL VALIDATION DIAGNOSTICS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Validation loss:\",\n    f\"{diagnostic_loss:.6f}\"\n)\n\nprint()\n\nprint(\n    diagnostic_df[\n        [\n            \"target\",\n            \"actual_positive\",\n            \"predicted_positive\",\n            \"TP\",\n            \"TN\",\n            \"FP\",\n            \"FN\",\n            \"sensitivity\",\n            \"specificity\",\n            \"precision\",\n            \"f1\",\n            \"mean_probability\"\n        ]\n    ].to_string(\n        index=False\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Aggregate diagnostics\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"AGGREGATE RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Mean F1:\",\n    f\"{diagnostic_df['f1'].mean():.4f}\"\n)\n\nprint(\n    \"Mean sensitivity:\",\n    f\"{diagnostic_df['sensitivity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean specificity:\",\n    f\"{diagnostic_df['specificity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean precision:\",\n    f\"{diagnostic_df['precision'].mean():.4f}\"\n)\n\nprint(\n    \"Total actual positives:\",\n    int(\n        diagnostic_df[\n            \"actual_positive\"\n        ].sum()\n    )\n)\n\nprint(\n    \"Total predicted positives:\",\n    int(\n        diagnostic_df[\n            \"predicted_positive\"\n        ].sum()\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Prediction-collapse checks\n# ---------------------------------------------------------------\n\nall_negative_targets = diagnostic_df[\n    diagnostic_df[\n        \"predicted_positive\"\n    ] == 0\n][\"target\"].tolist()\n\n\nall_positive_targets = diagnostic_df[\n    diagnostic_df[\n        \"predicted_positive\"\n    ] == len(\n        diagnostic_labels\n    )\n][\"target\"].tolist()\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"PREDICTION COLLAPSE CHECK\")\nprint(\"=\" * 70)\n\nprint(\n    \"All-negative targets:\",\n    all_negative_targets\n)\n\nprint(\n    \"All-positive targets:\",\n    all_positive_targets\n)\n\n\n# ---------------------------------------------------------------\n# Probability distribution\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PROBABILITY SUMMARY\")\nprint(\"=\" * 70)\n\nfor index, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    probabilities = (\n        diagnostic_probabilities[\n            :,\n            index\n        ]\n    )\n\n    print(\n        f\"{target:20s} \"\n        f\"mean={np.mean(probabilities):.4f} \"\n        f\"median={np.median(probabilities):.4f} \"\n        f\"p10={np.percentile(probabilities, 10):.4f} \"\n        f\"p90={np.percentile(probabilities, 90):.4f}\"\n    )\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 38 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved:\",\n    DIAGNOSTIC_PATH\n)\n\nprint()\nprint(\n    \"Use this diagnostic before continuing full training.\"\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 38 - FOLD 0 BEST MODEL DIAGNOSTIC\n# ================================================================\n\nprint(\"=\" * 70)\nprint(\"CELL 38 - FOLD 0 BEST MODEL DIAGNOSTIC\")\nprint(\"=\" * 70)\n\n\n# ---------------------------------------------------------------\n# Load best checkpoint\n# ---------------------------------------------------------------\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\n\nprint()\nprint(\"Best checkpoint epoch:\")\nprint(\n    checkpoint[\"epoch\"]\n)\n\nprint(\n    \"Best validation loss:\",\n    checkpoint[\"best_val_loss\"]\n)\n\n\n# ---------------------------------------------------------------\n# Run validation\n# ---------------------------------------------------------------\n\n(\n    diagnostic_loss,\n    diagnostic_labels,\n    diagnostic_probabilities,\n    diagnostic_predictions\n) = evaluate_fold0(\n    model,\n    val_loader_fold0,\n    criterion_fold,\n    DEVICE\n)\n\n\n# ---------------------------------------------------------------\n# Build detailed diagnostics\n# ---------------------------------------------------------------\n\ndiagnostic_rows = []\n\n\nfor index, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    y_true = diagnostic_labels[\n        :,\n        index\n    ]\n\n    y_prob = diagnostic_probabilities[\n        :,\n        index\n    ]\n\n    y_pred = diagnostic_predictions[\n        :,\n        index\n    ]\n\n\n    actual_positive = int(\n        np.sum(\n            y_true == 1\n        )\n    )\n\n    actual_negative = int(\n        np.sum(\n            y_true == 0\n        )\n    )\n\n    predicted_positive = int(\n        np.sum(\n            y_pred == 1\n        )\n    )\n\n    predicted_negative = int(\n        np.sum(\n            y_pred == 0\n        )\n    )\n\n\n    tp = int(\n        np.sum(\n            (y_true == 1)\n            &\n            (y_pred == 1)\n        )\n    )\n\n    tn = int(\n        np.sum(\n            (y_true == 0)\n            &\n            (y_pred == 0)\n        )\n    )\n\n    fp = int(\n        np.sum(\n            (y_true == 0)\n            &\n            (y_pred == 1)\n        )\n    )\n\n    fn = int(\n        np.sum(\n            (y_true == 1)\n            &\n            (y_pred == 0)\n        )\n    )\n\n\n    sensitivity = (\n        tp / (tp + fn)\n        if (tp + fn) > 0\n        else np.nan\n    )\n\n    specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else np.nan\n    )\n\n    precision = (\n        tp / (tp + fp)\n        if (tp + fp) > 0\n        else 0.0\n    )\n\n    f1 = (\n        2 * precision * sensitivity\n        / (precision + sensitivity)\n        if (\n            precision + sensitivity\n        ) > 0\n        else 0.0\n    )\n\n\n    diagnostic_rows.append(\n        {\n            \"target\": target,\n            \"actual_positive\": actual_positive,\n            \"actual_negative\": actual_negative,\n            \"predicted_positive\": predicted_positive,\n            \"predicted_negative\": predicted_negative,\n            \"TP\": tp,\n            \"TN\": tn,\n            \"FP\": fp,\n            \"FN\": fn,\n            \"mean_probability\": float(\n                np.mean(y_prob)\n            ),\n            \"min_probability\": float(\n                np.min(y_prob)\n            ),\n            \"max_probability\": float(\n                np.max(y_prob)\n            ),\n            \"sensitivity\": sensitivity,\n            \"specificity\": specificity,\n            \"precision\": precision,\n            \"f1\": f1\n        }\n    )\n\n\ndiagnostic_df = pd.DataFrame(\n    diagnostic_rows\n)\n\n\n# ---------------------------------------------------------------\n# Save diagnostic table\n# ---------------------------------------------------------------\n\nDIAGNOSTIC_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell38_fold0_best_model_diagnostics.csv\"\n)\n\ndiagnostic_df.to_csv(\n    DIAGNOSTIC_PATH,\n    index=False\n)\n\n\n# ---------------------------------------------------------------\n# Print diagnostics\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BEST MODEL VALIDATION DIAGNOSTICS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Validation loss:\",\n    f\"{diagnostic_loss:.6f}\"\n)\n\nprint()\n\nprint(\n    diagnostic_df[\n        [\n            \"target\",\n            \"actual_positive\",\n            \"predicted_positive\",\n            \"TP\",\n            \"TN\",\n            \"FP\",\n            \"FN\",\n            \"sensitivity\",\n            \"specificity\",\n            \"precision\",\n            \"f1\",\n            \"mean_probability\"\n        ]\n    ].to_string(\n        index=False\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Aggregate diagnostics\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"AGGREGATE RESULTS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Mean F1:\",\n    f\"{diagnostic_df['f1'].mean():.4f}\"\n)\n\nprint(\n    \"Mean sensitivity:\",\n    f\"{diagnostic_df['sensitivity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean specificity:\",\n    f\"{diagnostic_df['specificity'].mean():.4f}\"\n)\n\nprint(\n    \"Mean precision:\",\n    f\"{diagnostic_df['precision'].mean():.4f}\"\n)\n\nprint(\n    \"Total actual positives:\",\n    int(\n        diagnostic_df[\n            \"actual_positive\"\n        ].sum()\n    )\n)\n\nprint(\n    \"Total predicted positives:\",\n    int(\n        diagnostic_df[\n            \"predicted_positive\"\n        ].sum()\n    )\n)\n\n\n# ---------------------------------------------------------------\n# Prediction-collapse checks\n# ---------------------------------------------------------------\n\nall_negative_targets = diagnostic_df[\n    diagnostic_df[\n        \"predicted_positive\"\n    ] == 0\n][\"target\"].tolist()\n\n\nall_positive_targets = diagnostic_df[\n    diagnostic_df[\n        \"predicted_positive\"\n    ] == len(\n        diagnostic_labels\n    )\n][\"target\"].tolist()\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"PREDICTION COLLAPSE CHECK\")\nprint(\"=\" * 70)\n\nprint(\n    \"All-negative targets:\",\n    all_negative_targets\n)\n\nprint(\n    \"All-positive targets:\",\n    all_positive_targets\n)\n\n\n# ---------------------------------------------------------------\n# Probability distribution\n# ---------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PROBABILITY SUMMARY\")\nprint(\"=\" * 70)\n\nfor index, target in enumerate(\n    TARGET_COLUMNS\n):\n\n    probabilities = (\n        diagnostic_probabilities[\n            :,\n            index\n        ]\n    )\n\n    print(\n        f\"{target:20s} \"\n        f\"mean={np.mean(probabilities):.4f} \"\n        f\"median={np.median(probabilities):.4f} \"\n        f\"p10={np.percentile(probabilities, 10):.4f} \"\n        f\"p90={np.percentile(probabilities, 90):.4f}\"\n    )\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 38 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Saved:\",\n    DIAGNOSTIC_PATH\n)\n\nprint()\nprint(\n    \"Use this diagnostic before continuing full training.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:09:36.057689Z","iopub.execute_input":"2026-08-11T09:09:36.058509Z","iopub.status.idle":"2026-08-11T09:09:44.866422Z","shell.execute_reply.started":"2026-08-11T09:09:36.058471Z","shell.execute_reply":"2026-08-11T09:09:44.865314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 39 - FOLD 0 PROBABILITY SEPARATION AUDIT\n# ================================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score, f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 39 - FOLD 0 PROBABILITY SEPARATION AUDIT\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------\n# 1. VERIFY EXISTING NOTEBOOK STATE\n# ------------------------------------------------\n\nprint(\"\\nChecking existing notebook objects...\")\n\nprint(\"model:\", \"FOUND\" if \"model\" in globals() else \"MISSING\")\nprint(\"val_loader:\", \"FOUND\" if \"val_loader\" in globals() else \"MISSING\")\nprint(\"TARGETS:\", \"FOUND\" if \"TARGETS\" in globals() else \"MISSING\")\n\nif \"model\" not in globals():\n    raise RuntimeError(\"model is missing. Do not retrain. Send me this output.\")\n\nif \"val_loader\" not in globals():\n    raise RuntimeError(\"val_loader is missing. Do not retrain. Send me this output.\")\n\nif \"TARGETS\" not in globals():\n    raise RuntimeError(\"TARGETS is missing. Do not retrain. Send me this output.\")\n\n# ------------------------------------------------\n# 2. DETERMINE DEVICE FROM THE EXISTING MODEL\n# ------------------------------------------------\n\nDEVICE = next(model.parameters()).device\n\nprint(\"Device detected from model:\", DEVICE)\n\n# ------------------------------------------------\n# 3. LOAD THE BEST FOLD-0 CHECKPOINT\n# ------------------------------------------------\n\ncheckpoint_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell37_fold0_best_model.pt\"\n)\n\nprint(\"\\nLoading best Fold-0 checkpoint:\")\nprint(checkpoint_path)\n\ncheckpoint = torch.load(\n    checkpoint_path,\n    map_location=DEVICE,\n    weights_only=False\n)\n\nif isinstance(checkpoint, dict) and \"model_state_dict\" in checkpoint:\n    model.load_state_dict(checkpoint[\"model_state_dict\"])\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.to(DEVICE)\nmodel.eval()\n\nprint(\"Best checkpoint loaded: PASS\")\nprint(\"Model set to evaluation mode: PASS\")\n\n# ------------------------------------------------\n# 4. RUN VALIDATION INFERENCE\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"VALIDATION INFERENCE\")\nprint(\"=\" * 70)\n\nall_probabilities = []\nall_labels = []\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        images = batch[0].to(DEVICE)\n        labels = batch[1].to(DEVICE)\n\n        logits = model(images)\n\n        probs = torch.sigmoid(logits)\n\n        all_probabilities.append(\n            probs.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.cpu().numpy()\n        )\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\nlabels = np.concatenate(\n    all_labels,\n    axis=0\n)\n\nprint(\"Probability shape:\", probabilities.shape)\nprint(\"Label shape:\", labels.shape)\n\nif probabilities.shape != labels.shape:\n    raise RuntimeError(\n        f\"Shape mismatch: probabilities={probabilities.shape}, \"\n        f\"labels={labels.shape}\"\n    )\n\nif probabilities.shape[1] != len(TARGETS):\n    raise RuntimeError(\n        f\"Expected {len(TARGETS)} targets, \"\n        f\"received {probabilities.shape[1]} outputs.\"\n    )\n\nif not np.isfinite(probabilities).all():\n    raise RuntimeError(\"NaN/Inf detected in probabilities.\")\n\nprint(\"Inference: PASS\")\nprint(\"Shape check: PASS\")\nprint(\"NaN/Inf check: PASS\")\n\n# ------------------------------------------------\n# 5. PROBABILITY SEPARATION\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PROBABILITY SEPARATION\")\nprint(\"=\" * 70)\n\nseparation_results = []\n\nfor i, target in enumerate(TARGETS):\n\n    y_true = labels[:, i].astype(int)\n    y_prob = probabilities[:, i]\n\n    positive_probs = y_prob[y_true == 1]\n    negative_probs = y_prob[y_true == 0]\n\n    positive_count = len(positive_probs)\n    negative_count = len(negative_probs)\n\n    positive_mean = np.mean(positive_probs)\n    negative_mean = np.mean(negative_probs)\n\n    positive_median = np.median(positive_probs)\n    negative_median = np.median(negative_probs)\n\n    separation = (\n        positive_mean - negative_mean\n    )\n\n    auc = roc_auc_score(\n        y_true,\n        y_prob\n    )\n\n    separation_results.append({\n        \"target\": target,\n        \"positive_count\": positive_count,\n        \"negative_count\": negative_count,\n        \"positive_mean\": positive_mean,\n        \"negative_mean\": negative_mean,\n        \"positive_median\": positive_median,\n        \"negative_median\": negative_median,\n        \"probability_separation\": separation,\n        \"roc_auc\": auc\n    })\n\nseparation_df = pd.DataFrame(\n    separation_results\n)\n\nprint(\n    separation_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ------------------------------------------------\n# 6. THRESHOLD AUDIT\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PER-TARGET THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\nthresholds = np.arange(\n    0.20,\n    0.801,\n    0.025\n)\n\nthreshold_results = []\n\nfor i, target in enumerate(TARGETS):\n\n    y_true = labels[:, i].astype(int)\n    y_prob = probabilities[:, i]\n\n    baseline_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    baseline_f1 = f1_score(\n        y_true,\n        baseline_pred,\n        zero_division=0\n    )\n\n    best_threshold = 0.50\n    best_f1 = baseline_f1\n\n    for threshold in thresholds:\n\n        predictions = (\n            y_prob >= threshold\n        ).astype(int)\n\n        current_f1 = f1_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n\n        if current_f1 > best_f1:\n\n            best_f1 = current_f1\n            best_threshold = threshold\n\n    threshold_results.append({\n        \"target\": target,\n        \"f1_at_0.50\": baseline_f1,\n        \"best_threshold\": best_threshold,\n        \"best_f1\": best_f1,\n        \"f1_improvement\": (\n            best_f1 - baseline_f1\n        )\n    })\n\nthreshold_df = pd.DataFrame(\n    threshold_results\n)\n\nprint(\n    threshold_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ------------------------------------------------\n# 7. SUMMARY\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 39 SUMMARY\")\nprint(\"=\" * 70)\n\nmean_auc = separation_df[\"roc_auc\"].mean()\n\nmean_separation = (\n    separation_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = (\n    threshold_df[\"f1_at_0.50\"].mean()\n)\n\nmean_best_f1 = (\n    threshold_df[\"best_f1\"].mean()\n)\n\nprint(\n    f\"Mean ROC-AUC:                 {mean_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation:  {mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50:               {mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1:       {mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement:     \"\n    f\"{mean_best_f1 - mean_f1_050:.4f}\"\n)\n\nprint(\n    \"\\nTargets with ROC-AUC >= 0.65:\",\n    int(\n        (separation_df[\"roc_auc\"] >= 0.65).sum()\n    )\n)\n\nprint(\n    \"Targets with negative probability separation:\",\n    int(\n        (\n            separation_df[\n                \"probability_separation\"\n            ] < 0\n        ).sum()\n    )\n)\n\nprint(\n    \"Targets with >0.05 F1 improvement from threshold:\",\n    int(\n        (\n            threshold_df[\n                \"f1_improvement\"\n            ] > 0.05\n        ).sum()\n    )\n)\n\n# ------------------------------------------------\n# 8. SAVE RESULTS\n# ------------------------------------------------\n\noutput_dir = (\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nseparation_path = (\n    output_dir +\n    \"/cell39_probability_separation.csv\"\n)\n\nthreshold_path = (\n    output_dir +\n    \"/cell39_threshold_audit.csv\"\n)\n\nseparation_df.to_csv(\n    separation_path,\n    index=False\n)\n\nthreshold_df.to_csv(\n    threshold_path,\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(separation_path)\nprint(threshold_path)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 39 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:41:26.824517Z","iopub.execute_input":"2026-08-11T09:41:26.825132Z","iopub.status.idle":"2026-08-11T09:41:27.950269Z","shell.execute_reply.started":"2026-08-11T09:41:26.825094Z","shell.execute_reply":"2026-08-11T09:41:27.949051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 39 - PROBABILITY SEPARATION + THRESHOLD AUDIT\n# ================================================================\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score, f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 39 - PROBABILITY SEPARATION + THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------\n# 1. VERIFY EXISTING NOTEBOOK STATE\n# ------------------------------------------------\n\nprint(\"\\nChecking notebook state...\")\n\nrequired_objects = [\n    \"model\",\n    \"val_loader\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"model       : AVAILABLE\")\nprint(\"val_loader  : AVAILABLE\")\nprint(\"TARGETS     : AVAILABLE\")\n\n# ------------------------------------------------\n# 2. DETERMINE MODEL DEVICE\n# ------------------------------------------------\n\nMODEL_DEVICE = next(\n    model.parameters()\n).device\n\nprint(\"Model device:\", MODEL_DEVICE)\n\n# ------------------------------------------------\n# 3. LOAD BEST FOLD-0 CHECKPOINT\n# ------------------------------------------------\n\ncheckpoint_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell37_fold0_best_model.pt\"\n)\n\nprint(\"\\nLoading best Fold-0 checkpoint:\")\nprint(checkpoint_path)\n\ncheckpoint = torch.load(\n    checkpoint_path,\n    map_location=MODEL_DEVICE,\n    weights_only=False\n)\n\nif (\n    isinstance(checkpoint, dict)\n    and \"model_state_dict\" in checkpoint\n):\n    model.load_state_dict(\n        checkpoint[\"model_state_dict\"]\n    )\nelse:\n    model.load_state_dict(checkpoint)\n\nmodel.to(MODEL_DEVICE)\nmodel.eval()\n\nprint(\"Best Fold-0 checkpoint: LOADED\")\nprint(\"Model mode: evaluation\")\n\n# ------------------------------------------------\n# 4. INSPECT ACTUAL VAL_LOADER BATCH STRUCTURE\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INSPECTING VALIDATION BATCH STRUCTURE\")\nprint(\"=\" * 70)\n\nfirst_batch = next(iter(val_loader))\n\nprint(\"Batch type:\", type(first_batch))\n\nif isinstance(first_batch, dict):\n\n    print(\"Batch keys:\")\n    for key, value in first_batch.items():\n        if hasattr(value, \"shape\"):\n            print(\n                f\"  {key}: \"\n                f\"type={type(value).__name__}, \"\n                f\"shape={tuple(value.shape)}\"\n            )\n        else:\n            print(\n                f\"  {key}: \"\n                f\"type={type(value).__name__}\"\n            )\n\nelif isinstance(first_batch, (tuple, list)):\n\n    print(\n        \"Batch sequence length:\",\n        len(first_batch)\n    )\n\n    for i, value in enumerate(first_batch):\n\n        if hasattr(value, \"shape\"):\n            print(\n                f\"  [{i}]: \"\n                f\"type={type(value).__name__}, \"\n                f\"shape={tuple(value.shape)}\"\n            )\n        else:\n            print(\n                f\"  [{i}]: \"\n                f\"type={type(value).__name__}\"\n            )\n\nelse:\n\n    raise RuntimeError(\n        \"Unsupported validation batch type: \"\n        + str(type(first_batch))\n    )\n\n# ------------------------------------------------\n# 5. GENERIC BATCH EXTRACTION\n# ------------------------------------------------\n\ndef extract_image_and_label(batch):\n    \"\"\"\n    Extract image and label tensors from the existing\n    Dataset/DataLoader batch without assuming fixed\n    dictionary key names or tuple positions.\n    \"\"\"\n\n    candidates = []\n\n    if isinstance(batch, dict):\n\n        items = list(batch.items())\n\n        for key, value in items:\n\n            if torch.is_tensor(value):\n\n                candidates.append(\n                    (str(key), value)\n                )\n\n    elif isinstance(batch, (tuple, list)):\n\n        for idx, value in enumerate(batch):\n\n            if torch.is_tensor(value):\n\n                candidates.append(\n                    (str(idx), value)\n                )\n\n    else:\n\n        raise RuntimeError(\n            \"Unsupported batch type: \"\n            + str(type(batch))\n        )\n\n    image_candidates = []\n    label_candidates = []\n\n    for name, tensor in candidates:\n\n        if tensor.ndim == 4:\n            image_candidates.append(\n                (name, tensor)\n            )\n\n        elif tensor.ndim == 2:\n            if tensor.shape[-1] == len(TARGETS):\n                label_candidates.append(\n                    (name, tensor)\n                )\n\n    if len(image_candidates) == 0:\n\n        raise RuntimeError(\n            \"Could not identify image tensor. \"\n            \"Expected a 4D tensor such as \"\n            \"(batch, channels, height, width).\"\n        )\n\n    if len(label_candidates) == 0:\n\n        raise RuntimeError(\n            \"Could not identify label tensor. \"\n            f\"Expected a 2D tensor with \"\n            f\"{len(TARGETS)} targets.\"\n        )\n\n    image_name, images = image_candidates[0]\n    label_name, labels = label_candidates[0]\n\n    return (\n        images,\n        labels,\n        image_name,\n        label_name\n    )\n\nimages_test, labels_test, image_key, label_key = (\n    extract_image_and_label(first_batch)\n)\n\nprint(\"\\nIdentified image field:\", image_key)\nprint(\n    \"Identified label field:\",\n    label_key\n)\n\nprint(\n    \"Image batch shape:\",\n    tuple(images_test.shape)\n)\n\nprint(\n    \"Label batch shape:\",\n    tuple(labels_test.shape)\n)\n\n# ------------------------------------------------\n# 6. VALIDATE BATCH STRUCTURE\n# ------------------------------------------------\n\nif images_test.ndim != 4:\n    raise RuntimeError(\n        \"Image tensor is not 4-dimensional.\"\n    )\n\nif labels_test.ndim != 2:\n    raise RuntimeError(\n        \"Label tensor is not 2-dimensional.\"\n    )\n\nif labels_test.shape[1] != len(TARGETS):\n    raise RuntimeError(\n        f\"Label tensor has {labels_test.shape[1]} \"\n        f\"targets, expected {len(TARGETS)}.\"\n    )\n\nif images_test.shape[0] != labels_test.shape[0]:\n    raise RuntimeError(\n        \"Image and label batch sizes do not match.\"\n    )\n\nprint(\"Batch structure: PASS\")\n\n# ------------------------------------------------\n# 7. RUN VALIDATION INFERENCE\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"VALIDATION INFERENCE\")\nprint(\"=\" * 70)\n\nall_probabilities = []\nall_labels = []\n\nbatch_count = 0\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        images, labels, _, _ = (\n            extract_image_and_label(batch)\n        )\n\n        images = images.to(\n            MODEL_DEVICE,\n            non_blocking=True\n        )\n\n        labels = labels.to(\n            MODEL_DEVICE,\n            non_blocking=True\n        )\n\n        logits = model(images)\n\n        probabilities_batch = torch.sigmoid(\n            logits\n        )\n\n        all_probabilities.append(\n            probabilities_batch.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.cpu().numpy()\n        )\n\n        batch_count += 1\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\nlabels = np.concatenate(\n    all_labels,\n    axis=0\n)\n\nprint(\"Validation batches:\", batch_count)\nprint(\n    \"Probability shape:\",\n    probabilities.shape\n)\n\nprint(\n    \"Label shape:\",\n    labels.shape\n)\n\n# ------------------------------------------------\n# 8. FINAL INFERENCE VALIDATION\n# ------------------------------------------------\n\nexpected_shape = (\n    probabilities.shape[0],\n    len(TARGETS)\n)\n\nif probabilities.shape != expected_shape:\n    raise RuntimeError(\n        f\"Unexpected probability shape: \"\n        f\"{probabilities.shape}\"\n    )\n\nif labels.shape != expected_shape:\n    raise RuntimeError(\n        f\"Unexpected label shape: \"\n        f\"{labels.shape}\"\n    )\n\nif not np.isfinite(probabilities).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in probabilities.\"\n    )\n\nprint(\"Probability validity: PASS\")\nprint(\"Shape consistency: PASS\")\n\n# ------------------------------------------------\n# 9. PROBABILITY SEPARATION\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PROBABILITY SEPARATION\")\nprint(\"=\" * 70)\n\nseparation_rows = []\n\nfor idx, target in enumerate(TARGETS):\n\n    y_true = labels[:, idx].astype(int)\n    y_prob = probabilities[:, idx]\n\n    positive_probs = y_prob[y_true == 1]\n    negative_probs = y_prob[y_true == 0]\n\n    positive_count = len(\n        positive_probs\n    )\n\n    negative_count = len(\n        negative_probs\n    )\n\n    if (\n        positive_count > 0\n        and negative_count > 0\n    ):\n\n        positive_mean = float(\n            positive_probs.mean()\n        )\n\n        negative_mean = float(\n            negative_probs.mean()\n        )\n\n        positive_median = float(\n            np.median(positive_probs)\n        )\n\n        negative_median = float(\n            np.median(negative_probs)\n        )\n\n        separation = (\n            positive_mean\n            - negative_mean\n        )\n\n        auc = float(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n\n    else:\n\n        positive_mean = np.nan\n        negative_mean = np.nan\n        positive_median = np.nan\n        negative_median = np.nan\n        separation = np.nan\n        auc = np.nan\n\n    separation_rows.append({\n        \"target\": target,\n        \"positive_count\": positive_count,\n        \"negative_count\": negative_count,\n        \"positive_mean\": positive_mean,\n        \"negative_mean\": negative_mean,\n        \"positive_median\": positive_median,\n        \"negative_median\": negative_median,\n        \"probability_separation\": separation,\n        \"roc_auc\": auc\n    })\n\nseparation_df = pd.DataFrame(\n    separation_rows\n)\n\nprint(\n    separation_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ------------------------------------------------\n# 10. THRESHOLD AUDIT\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PER-TARGET THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\nthreshold_values = np.arange(\n    0.20,\n    0.801,\n    0.025\n)\n\nthreshold_rows = []\n\nfor idx, target in enumerate(TARGETS):\n\n    y_true = labels[:, idx].astype(int)\n    y_prob = probabilities[:, idx]\n\n    f1_at_050 = f1_score(\n        y_true,\n        (y_prob >= 0.50).astype(int),\n        zero_division=0\n    )\n\n    best_threshold = 0.50\n    best_f1 = f1_at_050\n\n    for threshold in threshold_values:\n\n        predictions = (\n            y_prob >= threshold\n        ).astype(int)\n\n        current_f1 = f1_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n\n        if current_f1 > best_f1:\n\n            best_f1 = current_f1\n            best_threshold = float(\n                threshold\n            )\n\n    threshold_rows.append({\n        \"target\": target,\n        \"f1_at_0.50\": float(\n            f1_at_050\n        ),\n        \"best_threshold\": float(\n            best_threshold\n        ),\n        \"best_f1\": float(\n            best_f1\n        ),\n        \"f1_improvement\": float(\n            best_f1 - f1_at_050\n        )\n    })\n\nthreshold_df = pd.DataFrame(\n    threshold_rows\n)\n\nprint(\n    threshold_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ------------------------------------------------\n# 11. OVERALL SUMMARY\n# ------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 39 SUMMARY\")\nprint(\"=\" * 70)\n\nmean_auc = float(\n    separation_df[\"roc_auc\"].mean()\n)\n\nmean_separation = float(\n    separation_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    threshold_df[\"f1_at_0.50\"].mean()\n)\n\nmean_best_f1 = float(\n    threshold_df[\"best_f1\"].mean()\n)\n\nprint(\n    f\"Mean ROC-AUC: \"\n    f\"{mean_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50: \"\n    f\"{mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{mean_best_f1 - mean_f1_050:.4f}\"\n)\n\nprint(\n    \"\\nTargets with ROC-AUC >= 0.65:\",\n    int(\n        (\n            separation_df[\"roc_auc\"]\n            >= 0.65\n        ).sum()\n    )\n)\n\nprint(\n    \"Targets with negative probability separation:\",\n    int(\n        (\n            separation_df[\n                \"probability_separation\"\n            ] < 0\n        ).sum()\n    )\n)\n\nprint(\n    \"Targets with >0.05 F1 improvement:\",\n    int(\n        (\n            threshold_df[\n                \"f1_improvement\"\n            ] > 0.05\n        ).sum()\n    )\n)\n\n# ------------------------------------------------\n# 12. SAVE RESULTS\n# ------------------------------------------------\n\noutput_dir = (\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\nseparation_path = (\n    output_dir +\n    \"/cell39_probability_separation.csv\"\n)\n\nthreshold_path = (\n    output_dir +\n    \"/cell39_threshold_audit.csv\"\n)\n\nseparation_df.to_csv(\n    separation_path,\n    index=False\n)\n\nthreshold_df.to_csv(\n    threshold_path,\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(separation_path)\nprint(threshold_path)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 39 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:44:15.229654Z","iopub.execute_input":"2026-08-11T09:44:15.230439Z","iopub.status.idle":"2026-08-11T09:44:26.175161Z","shell.execute_reply.started":"2026-08-11T09:44:15.230404Z","shell.execute_reply":"2026-08-11T09:44:26.173994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 40 - FOLD 1 TRAINING\n# ================================================================\n\nimport os\nimport copy\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom sklearn.metrics import f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 40 - FOLD 1 TRAINING\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nSEED = 42\nFOLD = 1\nMAX_EPOCHS = 10\nPATIENCE = 3\nLEARNING_RATE = 0.001\nWEIGHT_DECAY = 0.0001\nMAX_GRAD_NORM = 1.0\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\nCHECKPOINT_PATH = os.path.join(\n    AUDIT_DIR, \"cell40_fold1_best_model.pt\"\n)\nHISTORY_PATH = os.path.join(\n    AUDIT_DIR, \"cell40_fold1_history.csv\"\n)\nMETRICS_PATH = os.path.join(\n    AUDIT_DIR, \"cell40_fold1_metrics.csv\"\n)\n\n# ----------------------------------------------------------------\n# REPRODUCIBILITY\n# ----------------------------------------------------------------\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(f\"Fold: {FOLD}\")\nprint(f\"Maximum epochs: {MAX_EPOCHS}\")\nprint(f\"Early stopping patience: {PATIENCE}\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECT CHECK\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"train_loader\",\n    \"val_loader\",\n    \"TARGETS\",\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# VERIFY MODEL\n# ----------------------------------------------------------------\n\nmodel = model.to(device)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MODEL STATE\")\nprint(\"=\" * 70)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(\n    p.numel() for p in model.parameters()\n    if p.requires_grad\n)\n\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters before reset: {trainable_params:,}\")\n\nif total_params != 11_367_372:\n    print(\n        f\"WARNING: Expected 11,367,372 parameters from the \"\n        f\"validated Fold-0 model, found {total_params:,}.\"\n    )\n\nif trainable_params != 200_268:\n    print(\n        f\"WARNING: Expected 200,268 trainable parameters, \"\n        f\"found {trainable_params:,}.\"\n    )\n\n# ----------------------------------------------------------------\n# CREATE FRESH FOLD-1 MODEL\n#\n# Cell 39 loaded the best Fold-0 checkpoint into `model`.\n# The backbone/frozen parameters were never updated during Fold 0.\n#\n# We therefore copy the validated architecture and reset ONLY\n# parameters belonging to trainable modules.\n# ----------------------------------------------------------------\n\nfold1_model = copy.deepcopy(model)\n\nreset_parameter_count = 0\nreset_module_names = []\n\nfor module_name, module in fold1_model.named_modules():\n\n    if not any(\n        parameter.requires_grad\n        for parameter in module.parameters(recurse=False)\n    ):\n        continue\n\n    if hasattr(module, \"reset_parameters\"):\n        module.reset_parameters()\n\n        for parameter in module.parameters(recurse=False):\n            if parameter.requires_grad:\n                reset_parameter_count += parameter.numel()\n\n        reset_module_names.append(module_name)\n\nfold1_model = fold1_model.to(device)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"RESETTING MODEL FOR FOLD 1\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Trainable parameters after reset: \"\n    f\"{sum(p.numel() for p in fold1_model.parameters() if p.requires_grad):,}\"\n)\n\nprint(f\"Trainable modules reset: {len(reset_module_names)}\")\nprint(f\"Parameters reset: {reset_parameter_count:,}\")\n\n# ----------------------------------------------------------------\n# VERIFY TRAINABLE PARAMETER COUNT\n# ----------------------------------------------------------------\n\nfold1_trainable_params = sum(\n    p.numel()\n    for p in fold1_model.parameters()\n    if p.requires_grad\n)\n\nif fold1_trainable_params != trainable_params:\n    raise RuntimeError(\n        \"Trainable parameter count changed during reset: \"\n        f\"{trainable_params:,} -> {fold1_trainable_params:,}\"\n    )\n\nprint(\"Trainable parameter count: PASS\")\n\n# ----------------------------------------------------------------\n# CLASS WEIGHTS\n#\n# Recalculate weights strictly from the Fold-1 training studies.\n# This prevents validation labels from influencing the loss.\n# ----------------------------------------------------------------\n\nif \"modeling_df\" in globals():\n    fold_table = modeling_df.copy()\nelif \"modeling_table\" in globals():\n    fold_table = modeling_table.copy()\nelse:\n    fold_table = pd.read_csv(\n        os.path.join(\n            AUDIT_DIR,\n            \"cell29_modeling_studies.csv\"\n        )\n    )\n\nif \"fold_table\" not in globals():\n    raise RuntimeError(\"Unable to load modeling table.\")\n\nif \"fold\" not in fold_table.columns:\n    raise RuntimeError(\n        \"Modeling table does not contain the required 'fold' column.\"\n    )\n\ntrain_fold_df = fold_table[\n    fold_table[\"fold\"] != FOLD\n].copy()\n\nval_fold_df = fold_table[\n    fold_table[\"fold\"] == FOLD\n].copy()\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD DATA\")\nprint(\"=\" * 70)\n\nprint(f\"Training studies: {len(train_fold_df)}\")\nprint(f\"Validation studies: {len(val_fold_df)}\")\n\nif len(train_fold_df) == 0 or len(val_fold_df) == 0:\n    raise RuntimeError(\"Fold 1 has empty training or validation data.\")\n\n# ----------------------------------------------------------------\n# POSITIVE WEIGHTS\n# ----------------------------------------------------------------\n\npositive_weights = []\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 1 CLASS WEIGHTS\")\nprint(\"=\" * 70)\n\nfor target in TARGETS:\n\n    if target not in train_fold_df.columns:\n        raise RuntimeError(\n            f\"Target column missing from modeling table: {target}\"\n        )\n\n    positives = int(train_fold_df[target].sum())\n    negatives = int(len(train_fold_df) - positives)\n\n    if positives == 0:\n        raise RuntimeError(\n            f\"Fold 1 training set has zero positive samples for {target}.\"\n        )\n\n    if negatives == 0:\n        raise RuntimeError(\n            f\"Fold 1 training set has zero negative samples for {target}.\"\n        )\n\n    pos_weight = negatives / positives\n    positive_weights.append(pos_weight)\n\n    print(\n        f\"{target:22s} \"\n        f\"positive={positives:2d} \"\n        f\"negative={negatives:2d} \"\n        f\"pos_weight={pos_weight:.4f}\"\n    )\n\npos_weight_tensor = torch.tensor(\n    positive_weights,\n    dtype=torch.float32,\n    device=device\n)\n\n# ----------------------------------------------------------------\n# LOSS\n# ----------------------------------------------------------------\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_tensor\n)\n\n# ----------------------------------------------------------------\n# OPTIMIZER\n# ----------------------------------------------------------------\n\noptimizer = torch.optim.AdamW(\n    [\n        parameter\n        for parameter in fold1_model.parameters()\n        if parameter.requires_grad\n    ],\n    lr=LEARNING_RATE,\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"min\",\n    factor=0.5,\n    patience=1\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINING CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(f\"Loss: BCEWithLogitsLoss\")\nprint(f\"Optimizer: AdamW\")\nprint(f\"Learning rate: {LEARNING_RATE}\")\nprint(f\"Weight decay: {WEIGHT_DECAY}\")\nprint(f\"Gradient clipping: {MAX_GRAD_NORM}\")\nprint(f\"Scheduler: ReduceLROnPlateau\")\n\n# ----------------------------------------------------------------\n# TRAINING / VALIDATION FUNCTIONS\n# ----------------------------------------------------------------\n\ndef run_training_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n\n    running_loss = 0.0\n    batches = 0\n\n    for batch in loader:\n\n        images = batch[\"image\"].to(device, non_blocking=True)\n        labels = batch[\"label\"].to(device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        logits = model(images)\n\n        loss = criterion(logits, labels)\n\n        if not torch.isfinite(loss):\n            raise RuntimeError(\n                \"Non-finite training loss detected.\"\n            )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            MAX_GRAD_NORM\n        )\n\n        optimizer.step()\n\n        running_loss += loss.item()\n        batches += 1\n\n    if batches == 0:\n        raise RuntimeError(\"Training loader produced zero batches.\")\n\n    return running_loss / batches\n\n\n@torch.no_grad()\ndef run_validation_epoch(model, loader, criterion, device):\n    model.eval()\n\n    running_loss = 0.0\n    batches = 0\n\n    all_probabilities = []\n    all_labels = []\n\n    for batch in loader:\n\n        images = batch[\"image\"].to(device, non_blocking=True)\n        labels = batch[\"label\"].to(device, non_blocking=True)\n\n        logits = model(images)\n\n        loss = criterion(logits, labels)\n\n        if not torch.isfinite(loss):\n            raise RuntimeError(\n                \"Non-finite validation loss detected.\"\n            )\n\n        probabilities = torch.sigmoid(logits)\n\n        running_loss += loss.item()\n        batches += 1\n\n        all_probabilities.append(\n            probabilities.cpu()\n        )\n        all_labels.append(\n            labels.cpu()\n        )\n\n    if batches == 0:\n        raise RuntimeError(\"Validation loader produced zero batches.\")\n\n    probabilities = torch.cat(all_probabilities, dim=0).numpy()\n    labels = torch.cat(all_labels, dim=0).numpy()\n\n    predictions = (probabilities >= 0.5).astype(np.int32)\n\n    per_target_f1 = []\n\n    for target_index in range(len(TARGETS)):\n        f1 = f1_score(\n            labels[:, target_index],\n            predictions[:, target_index],\n            zero_division=0\n        )\n        per_target_f1.append(f1)\n\n    mean_f1 = float(np.mean(per_target_f1))\n\n    return (\n        running_loss / batches,\n        mean_f1,\n        probabilities,\n        labels\n    )\n\n# ----------------------------------------------------------------\n# TRAINING\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING FOLD 1 TRAINING\")\nprint(\"=\" * 70)\n\nhistory = []\n\nbest_val_loss = float(\"inf\")\nbest_epoch = 0\nepochs_without_improvement = 0\n\nfor epoch in range(1, MAX_EPOCHS + 1):\n\n    train_loss = run_training_epoch(\n        fold1_model,\n        train_loader,\n        criterion,\n        optimizer,\n        device\n    )\n\n    (\n        val_loss,\n        val_f1,\n        probabilities,\n        labels\n    ) = run_validation_epoch(\n        fold1_model,\n        val_loader,\n        criterion,\n        device\n    )\n\n    scheduler.step(val_loss)\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n\n    history.append({\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_loss,\n        \"mean_f1\": val_f1,\n        \"learning_rate\": current_lr\n    })\n\n    improved = val_loss < best_val_loss\n\n    if improved:\n\n        best_val_loss = val_loss\n        best_epoch = epoch\n        epochs_without_improvement = 0\n\n        torch.save(\n            {\n                \"model_state_dict\": fold1_model.state_dict(),\n                \"fold\": FOLD,\n                \"epoch\": epoch,\n                \"val_loss\": val_loss,\n                \"targets\": list(TARGETS),\n                \"input_channels\": 21,\n                \"image_size\": 224,\n                \"slices_per_plane\": 7,\n                \"planes\": [\n                    \"Sagittal\",\n                    \"Coronal\",\n                    \"Axial\"\n                ]\n            },\n            CHECKPOINT_PATH\n        )\n\n    else:\n        epochs_without_improvement += 1\n\n    print(f\"\\nEpoch {epoch:02d}/{MAX_EPOCHS}\")\n    print(f\"Train Loss: {train_loss:.5f}\")\n    print(f\"Val Loss:   {val_loss:.5f}\")\n    print(f\"Mean F1:    {val_f1:.4f}\")\n    print(f\"Learning Rate: {current_lr:.6f}\")\n    print(f\"Best Epoch: {best_epoch}\")\n    print(f\"Best Val Loss: {best_val_loss:.5f}\")\n\n    if epochs_without_improvement >= PATIENCE:\n        print(\n            f\"\\nEarly stopping triggered after \"\n            f\"{PATIENCE} epochs without improvement.\"\n        )\n        break\n\n# ----------------------------------------------------------------\n# RESTORE BEST CHECKPOINT\n# ----------------------------------------------------------------\n\nif not os.path.exists(CHECKPOINT_PATH):\n    raise RuntimeError(\n        \"Best Fold-1 checkpoint was not created.\"\n    )\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=device\n)\n\nfold1_model.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nfold1_model.eval()\n\n# ----------------------------------------------------------------\n# FINAL BEST-MODEL VALIDATION\n# ----------------------------------------------------------------\n\n(\n    final_val_loss,\n    final_mean_f1,\n    final_probabilities,\n    final_labels\n) = run_validation_epoch(\n    fold1_model,\n    val_loader,\n    criterion,\n    device\n)\n\n# ----------------------------------------------------------------\n# PER-TARGET METRICS\n# ----------------------------------------------------------------\n\nmetrics_rows = []\n\nfor target_index, target in enumerate(TARGETS):\n\n    y_true = final_labels[:, target_index]\n    y_prob = final_probabilities[:, target_index]\n    y_pred = (y_prob >= 0.5).astype(np.int32)\n\n    tp = int(((y_true == 1) & (y_pred == 1)).sum())\n    tn = int(((y_true == 0) & (y_pred == 0)).sum())\n    fp = int(((y_true == 0) & (y_pred == 1)).sum())\n    fn = int(((y_true == 1) & (y_pred == 0)).sum())\n\n    sensitivity = (\n        tp / (tp + fn)\n        if (tp + fn) > 0\n        else 0.0\n    )\n\n    specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else 0.0\n    )\n\n    precision = (\n        tp / (tp + fp)\n        if (tp + fp) > 0\n        else 0.0\n    )\n\n    f1 = (\n        2 * precision * sensitivity /\n        (precision + sensitivity)\n        if (precision + sensitivity) > 0\n        else 0.0\n    )\n\n    metrics_rows.append({\n        \"target\": target,\n        \"positive_count\": int(y_true.sum()),\n        \"negative_count\": int(len(y_true) - y_true.sum()),\n        \"predicted_positive\": int(y_pred.sum()),\n        \"TP\": tp,\n        \"TN\": tn,\n        \"FP\": fp,\n        \"FN\": fn,\n        \"sensitivity\": sensitivity,\n        \"specificity\": specificity,\n        \"precision\": precision,\n        \"f1\": f1,\n        \"mean_probability\": float(y_prob.mean())\n    })\n\nmetrics_df = pd.DataFrame(metrics_rows)\n\n# ----------------------------------------------------------------\n# SAVE HISTORY AND METRICS\n# ----------------------------------------------------------------\n\nhistory_df = pd.DataFrame(history)\n\nhistory_df.to_csv(\n    HISTORY_PATH,\n    index=False\n)\n\nmetrics_df.to_csv(\n    METRICS_PATH,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL OUTPUT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 1 TRAINING COMPLETE\")\nprint(\"=\" * 70)\n\nprint(f\"Best epoch: {best_epoch}\")\nprint(f\"Best validation loss: {best_val_loss:.6f}\")\nprint(f\"Final validation loss: {final_val_loss:.6f}\")\n\nprint(\"\\nPer-target validation metrics:\")\nprint(\n    metrics_df[\n        [\n            \"target\",\n            \"positive_count\",\n            \"negative_count\",\n            \"predicted_positive\",\n            \"sensitivity\",\n            \"specificity\",\n            \"precision\",\n            \"f1\"\n        ]\n    ].to_string(index=False)\n)\n\nprint(f\"\\nMean F1: {metrics_df['f1'].mean():.4f}\")\nprint(f\"Mean sensitivity: {metrics_df['sensitivity'].mean():.4f}\")\nprint(f\"Mean specificity: {metrics_df['specificity'].mean():.4f}\")\nprint(f\"Mean precision: {metrics_df['precision'].mean():.4f}\")\n\nprint(f\"\\nCheckpoint saved: {CHECKPOINT_PATH}\")\nprint(f\"History saved: {HISTORY_PATH}\")\nprint(f\"Metrics saved: {METRICS_PATH}\")\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 40 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T09:50:09.17935Z","iopub.execute_input":"2026-08-11T09:50:09.179673Z","iopub.status.idle":"2026-08-11T09:52:20.831026Z","shell.execute_reply.started":"2026-08-11T09:50:09.179647Z","shell.execute_reply":"2026-08-11T09:52:20.829974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 41 - FOLD 1 PROBABILITY SEPARATION + THRESHOLD AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score, f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 41 - FOLD 1 PROBABILITY SEPARATION + THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nFOLD = 1\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nCHECKPOINT_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell40_fold1_best_model.pt\"\n)\n\nSEPARATION_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell41_fold1_probability_separation.csv\"\n)\n\nTHRESHOLD_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell41_fold1_threshold_audit.csv\"\n)\n\n# ----------------------------------------------------------------\n# CHECK REQUIRED OBJECTS\n# ----------------------------------------------------------------\n\nprint(\"\\nChecking notebook state...\")\n\nrequired_objects = [\n    \"fold1_model\",\n    \"val_loader\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"fold1_model    : AVAILABLE\")\nprint(\"val_loader     : AVAILABLE\")\nprint(\"TARGETS        : AVAILABLE\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = next(fold1_model.parameters()).device\n\nprint(f\"Model device: {device}\")\n\n# ----------------------------------------------------------------\n# CHECKPOINT\n# ----------------------------------------------------------------\n\nif not os.path.exists(CHECKPOINT_PATH):\n    raise RuntimeError(\n        f\"Fold-1 checkpoint not found: {CHECKPOINT_PATH}\"\n    )\n\nprint(\"\\nLoading best Fold-1 checkpoint:\")\nprint(CHECKPOINT_PATH)\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=device\n)\n\nfold1_model.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nfold1_model.eval()\n\nprint(\"Best Fold-1 checkpoint: LOADED\")\nprint(\"Model mode: evaluation\")\n\n# ----------------------------------------------------------------\n# COLLECT VALIDATION PREDICTIONS\n# ----------------------------------------------------------------\n\nall_probabilities = []\nall_labels = []\n\nbatch_count = 0\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        if not isinstance(batch, dict):\n            raise RuntimeError(\n                f\"Unexpected batch type: {type(batch)}\"\n            )\n\n        if \"image\" not in batch or \"label\" not in batch:\n            raise RuntimeError(\n                \"Validation batch must contain 'image' and 'label'.\"\n            )\n\n        images = batch[\"image\"].to(device)\n        labels = batch[\"label\"].to(device)\n\n        logits = fold1_model(images)\n        probabilities = torch.sigmoid(logits)\n\n        if not torch.isfinite(probabilities).all():\n            raise RuntimeError(\n                \"NaN/Inf detected in validation probabilities.\"\n            )\n\n        all_probabilities.append(\n            probabilities.cpu()\n        )\n\n        all_labels.append(\n            labels.cpu()\n        )\n\n        batch_count += 1\n\nif batch_count == 0:\n    raise RuntimeError(\n        \"Validation loader produced zero batches.\"\n    )\n\nprobabilities = torch.cat(\n    all_probabilities,\n    dim=0\n).numpy()\n\nlabels = torch.cat(\n    all_labels,\n    dim=0\n).numpy()\n\nprint(\"\\nValidation batches:\", batch_count)\nprint(\"Probability shape:\", probabilities.shape)\nprint(\"Label shape:\", labels.shape)\n\nif probabilities.shape != labels.shape:\n    raise RuntimeError(\n        \"Probability and label shapes do not match.\"\n    )\n\nif probabilities.shape[1] != len(TARGETS):\n    raise RuntimeError(\n        \"Number of prediction columns does not match TARGETS.\"\n    )\n\nprint(\"Probability validity: PASS\")\nprint(\"Shape consistency: PASS\")\n\n# ----------------------------------------------------------------\n# PROBABILITY SEPARATION + ROC-AUC\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PROBABILITY SEPARATION\")\nprint(\"=\" * 70)\n\nseparation_rows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    positive_probabilities = y_prob[y_true == 1]\n    negative_probabilities = y_prob[y_true == 0]\n\n    positive_count = len(positive_probabilities)\n    negative_count = len(negative_probabilities)\n\n    if positive_count == 0 or negative_count == 0:\n        roc_auc = np.nan\n        probability_separation = np.nan\n        positive_mean = np.nan\n        negative_mean = np.nan\n        positive_median = np.nan\n        negative_median = np.nan\n    else:\n\n        positive_mean = float(\n            positive_probabilities.mean()\n        )\n\n        negative_mean = float(\n            negative_probabilities.mean()\n        )\n\n        positive_median = float(\n            np.median(positive_probabilities)\n        )\n\n        negative_median = float(\n            np.median(negative_probabilities)\n        )\n\n        probability_separation = (\n            positive_mean - negative_mean\n        )\n\n        roc_auc = float(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n\n    separation_rows.append({\n        \"target\": target,\n        \"positive_count\": positive_count,\n        \"negative_count\": negative_count,\n        \"positive_mean\": positive_mean,\n        \"negative_mean\": negative_mean,\n        \"positive_median\": positive_median,\n        \"negative_median\": negative_median,\n        \"probability_separation\": probability_separation,\n        \"roc_auc\": roc_auc\n    })\n\nseparation_df = pd.DataFrame(\n    separation_rows\n)\n\nprint(\n    separation_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# THRESHOLD AUDIT\n#\n# This is diagnostic only.\n# We DO NOT permanently adopt these thresholds.\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\nthresholds = np.arange(\n    0.20,\n    0.801,\n    0.025\n)\n\nthreshold_rows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    f1_at_050 = f1_score(\n        y_true,\n        (y_prob >= 0.50).astype(np.int32),\n        zero_division=0\n    )\n\n    best_f1 = -1.0\n    best_threshold = 0.50\n\n    for threshold in thresholds:\n\n        y_pred = (\n            y_prob >= threshold\n        ).astype(np.int32)\n\n        score = f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n\n        if score > best_f1:\n            best_f1 = float(score)\n            best_threshold = float(threshold)\n\n    threshold_rows.append({\n        \"target\": target,\n        \"f1_at_0.50\": float(f1_at_050),\n        \"best_threshold\": best_threshold,\n        \"best_f1\": best_f1,\n        \"f1_improvement\": best_f1 - float(f1_at_050)\n    })\n\nthreshold_df = pd.DataFrame(\n    threshold_rows\n)\n\nprint(\n    threshold_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------\n\nmean_roc_auc = float(\n    separation_df[\"roc_auc\"].mean()\n)\n\nmean_separation = float(\n    separation_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    threshold_df[\"f1_at_0.50\"].mean()\n)\n\nmean_best_f1 = float(\n    threshold_df[\"best_f1\"].mean()\n)\n\npotential_improvement = (\n    mean_best_f1 - mean_f1_050\n)\n\nnegative_separation_targets = (\n    separation_df[\n        separation_df[\"probability_separation\"] < 0\n    ][\"target\"].tolist()\n)\n\nhigh_auc_targets = (\n    separation_df[\n        separation_df[\"roc_auc\"] >= 0.65\n    ][\"target\"].tolist()\n)\n\nlarge_threshold_improvement_targets = (\n    threshold_df[\n        threshold_df[\"f1_improvement\"] > 0.05\n    ][\"target\"].tolist()\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 1 DIAGNOSTIC SUMMARY\")\nprint(\"=\" * 70)\n\nprint(f\"Mean ROC-AUC: {mean_roc_auc:.4f}\")\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\nprint(f\"Mean F1 @ 0.50: {mean_f1_050:.4f}\")\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{potential_improvement:.4f}\"\n)\n\nprint(\n    f\"\\nTargets with ROC-AUC >= 0.65: \"\n    f\"{len(high_auc_targets)}\"\n)\n\nif high_auc_targets:\n    print(high_auc_targets)\n\nprint(\n    f\"\\nTargets with negative probability separation: \"\n    f\"{len(negative_separation_targets)}\"\n)\n\nif negative_separation_targets:\n    print(negative_separation_targets)\n\nprint(\n    f\"\\nTargets with >0.05 F1 improvement: \"\n    f\"{len(large_threshold_improvement_targets)}\"\n)\n\nif large_threshold_improvement_targets:\n    print(large_threshold_improvement_targets)\n\n# ----------------------------------------------------------------\n# SAVE RESULTS\n# ----------------------------------------------------------------\n\nseparation_df.to_csv(\n    SEPARATION_PATH,\n    index=False\n)\n\nthreshold_df.to_csv(\n    THRESHOLD_PATH,\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(SEPARATION_PATH)\nprint(THRESHOLD_PATH)\n\n# ----------------------------------------------------------------\n# COMPLETION\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 41 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:01:26.936473Z","iopub.execute_input":"2026-08-11T10:01:26.936871Z","iopub.status.idle":"2026-08-11T10:01:38.318813Z","shell.execute_reply.started":"2026-08-11T10:01:26.936838Z","shell.execute_reply":"2026-08-11T10:01:38.317858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 42 - FOLD 2 DATA PREPARATION\n# ================================================================\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 70)\nprint(\"CELL 42 - FOLD 2 DATA PREPARATION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODELING_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nFOLDS_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_study_folds.csv\"\n)\n\nFOLD = 2\nBATCH_SIZE = 2\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"KneeStudyDataset\",\n    \"TARGETS\",\n    \"model\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required dataset/model objects: PASS\")\n\n# ----------------------------------------------------------------\n# LOAD EXACT SAVED TABLES\n# ----------------------------------------------------------------\n\nif not os.path.exists(MODELING_PATH):\n    raise RuntimeError(\n        f\"Missing modeling table: {MODELING_PATH}\"\n    )\n\nif not os.path.exists(FOLDS_PATH):\n    raise RuntimeError(\n        f\"Missing fold table: {FOLDS_PATH}\"\n    )\n\nmodeling_table = pd.read_csv(\n    MODELING_PATH\n)\n\nstudy_folds = pd.read_csv(\n    FOLDS_PATH\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"RESTORED MODELING DATA\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Modeling table shape: \"\n    f\"{modeling_table.shape}\"\n)\n\nprint(\n    f\"Fold table shape: \"\n    f\"{study_folds.shape}\"\n)\n\n# ----------------------------------------------------------------\n# EXACT SCHEMA INSPECTION\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 29 MODELING TABLE COLUMNS\")\nprint(\"=\" * 70)\n\nfor index, column in enumerate(\n    modeling_table.columns\n):\n    print(\n        f\"{index:02d}: {column}\"\n    )\n\n# ----------------------------------------------------------------\n# TARGET VALIDATION\n# ----------------------------------------------------------------\n\nmissing_targets = [\n    target\n    for target in TARGETS\n    if target not in modeling_table.columns\n]\n\nif missing_targets:\n    raise RuntimeError(\n        \"Missing target columns: \"\n        + \", \".join(missing_targets)\n    )\n\nprint(\"\\nTarget columns: PASS\")\n\n# ----------------------------------------------------------------\n# STUDY UID VALIDATION\n# ----------------------------------------------------------------\n\nif \"StudyInstanceUID\" not in modeling_table.columns:\n    raise RuntimeError(\n        \"StudyInstanceUID is missing from modeling table.\"\n    )\n\nif \"StudyInstanceUID\" not in study_folds.columns:\n    raise RuntimeError(\n        \"StudyInstanceUID is missing from fold table.\"\n    )\n\nif \"fold\" not in study_folds.columns:\n    raise RuntimeError(\n        \"fold column is missing from fold table.\"\n    )\n\n# ----------------------------------------------------------------\n# RESTORE FOLD ASSIGNMENT\n# ----------------------------------------------------------------\n\nstudy_folds[\"fold\"] = (\n    study_folds[\"fold\"]\n    .astype(int)\n)\n\nfold2_validation_ids = set(\n    study_folds.loc[\n        study_folds[\"fold\"] == FOLD,\n        \"StudyInstanceUID\"\n    ]\n)\n\nfold2_training_ids = set(\n    study_folds.loc[\n        study_folds[\"fold\"] != FOLD,\n        \"StudyInstanceUID\"\n    ]\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 STUDY SPLIT\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Training studies: \"\n    f\"{len(fold2_training_ids)}\"\n)\n\nprint(\n    f\"Validation studies: \"\n    f\"{len(fold2_validation_ids)}\"\n)\n\noverlap = (\n    fold2_training_ids\n    .intersection(\n        fold2_validation_ids\n    )\n)\n\nprint(\n    f\"Train/validation overlap: \"\n    f\"{len(overlap)}\"\n)\n\nif overlap:\n    raise RuntimeError(\n        \"Fold-2 train/validation study leakage detected.\"\n    )\n\nif len(fold2_training_ids) != 39:\n    raise RuntimeError(\n        \"Expected 39 Fold-2 training studies.\"\n    )\n\nif len(fold2_validation_ids) != 19:\n    raise RuntimeError(\n        \"Expected 19 Fold-2 validation studies.\"\n    )\n\nprint(\"Fold-2 split: PASS\")\n\n# ----------------------------------------------------------------\n# BUILD EXACT FOLD TABLES\n# ----------------------------------------------------------------\n\nfold2_train_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].isin(\n        fold2_training_ids\n    )\n].copy()\n\nfold2_val_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].isin(\n        fold2_validation_ids\n    )\n].copy()\n\nif len(fold2_train_table) != 39:\n    raise RuntimeError(\n        \"Fold-2 training table does not contain 39 rows.\"\n    )\n\nif len(fold2_val_table) != 19:\n    raise RuntimeError(\n        \"Fold-2 validation table does not contain 19 rows.\"\n    )\n\n# ----------------------------------------------------------------\n# IMPORTANT:\n# DO NOT GUESS PRIMARY SERIES COLUMN NAMES.\n#\n# Instead, identify the actual columns from Cell 29.\n# ----------------------------------------------------------------\n\nseries_uid_columns = [\n    column\n    for column in modeling_table.columns\n    if (\n        \"SeriesInstanceUID\" in column\n        or \"series\" in column.lower()\n        and \"uid\" in column.lower()\n    )\n]\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SERIES UID COLUMNS FOUND\")\nprint(\"=\" * 70)\n\nfor column in series_uid_columns:\n    print(column)\n\nif len(series_uid_columns) == 0:\n    print(\n        \"No series UID columns detected by name.\"\n    )\n\n# ----------------------------------------------------------------\n# EXISTING DATASET CLASS CHECK\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"EXISTING DATASET CLASS\")\nprint(\"=\" * 70)\n\nprint(\n    f\"KneeStudyDataset: \"\n    f\"{KneeStudyDataset}\"\n)\n\n# ----------------------------------------------------------------\n# DO NOT CREATE DATASETS YET.\n#\n# We first verify the exact Cell-29 schema. This prevents another\n# fabricated column-name error and preserves the Cell-33 pipeline.\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 42 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Modeling table restored: PASS\"\n)\n\nprint(\n    \"Fold table restored: PASS\"\n)\n\nprint(\n    \"Fold-2 split: PASS\"\n)\n\nprint(\n    \"Exact modeling-table schema printed: PASS\"\n)\n\nprint(\n    \"No fabricated primary-series columns used.\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No dataset pipeline modified.\"\n)\n\nprint(\"\\nNext step requires the actual column names printed above.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:12:12.899139Z","iopub.execute_input":"2026-08-11T10:12:12.899518Z","iopub.status.idle":"2026-08-11T10:12:12.933144Z","shell.execute_reply.started":"2026-08-11T10:12:12.899485Z","shell.execute_reply":"2026-08-11T10:12:12.931788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 43 - FOLD 2 DATASET + DATALOADER\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 70)\nprint(\"CELL 43 - FOLD 2 DATASET + DATALOADER\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nFOLD = 2\nBATCH_SIZE = 2\n\nIMAGE_SIZE = 224\nSLICES_PER_PLANE = 7\n\nEXPECTED_CHANNELS = 21\n\nPLANE_COLUMNS = {\n    \"Sagittal\": \"Sagittal_SeriesInstanceUID\",\n    \"Coronal\": \"Coronal_SeriesInstanceUID\",\n    \"Axial\": \"Axial_SeriesInstanceUID\"\n}\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"modeling_table\",\n    \"study_folds\",\n    \"KneeStudyDataset\",\n    \"TARGETS\",\n    \"model\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# VERIFY EXACT MODELING TABLE SCHEMA\n# ----------------------------------------------------------------\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\",\n    \"fold\"\n] + list(TARGETS)\n\nmissing_columns = [\n    column\n    for column in required_columns\n    if column not in modeling_table.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Missing required modeling-table columns: \"\n        + \", \".join(missing_columns)\n    )\n\nprint(\"Modeling-table schema: PASS\")\n\n# ----------------------------------------------------------------\n# VERIFY PRIMARY SERIES COMPLETENESS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PRIMARY SERIES VALIDATION\")\nprint(\"=\" * 70)\n\nfor plane, column in PLANE_COLUMNS.items():\n\n    missing_count = (\n        modeling_table[column]\n        .isna()\n        .sum()\n    )\n\n    empty_count = (\n        modeling_table[column]\n        .astype(str)\n        .str.strip()\n        .eq(\"\")\n        .sum()\n    )\n\n    print(\n        f\"{plane:8s}: \"\n        f\"missing={missing_count}, \"\n        f\"empty={empty_count}, \"\n        f\"unique={modeling_table[column].nunique()}\"\n    )\n\n    if missing_count > 0 or empty_count > 0:\n        raise RuntimeError(\n            f\"{plane} primary series contains missing values.\"\n        )\n\nprint(\"All three primary planes: PASS\")\n\n# ----------------------------------------------------------------\n# RESTORE FOLD-2 STUDY IDS\n# ----------------------------------------------------------------\n\nfold2_validation_ids = set(\n    study_folds.loc[\n        study_folds[\"fold\"].astype(int) == FOLD,\n        \"StudyInstanceUID\"\n    ]\n)\n\nfold2_training_ids = set(\n    study_folds.loc[\n        study_folds[\"fold\"].astype(int) != FOLD,\n        \"StudyInstanceUID\"\n    ]\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 STUDY SPLIT\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Training studies: \"\n    f\"{len(fold2_training_ids)}\"\n)\n\nprint(\n    f\"Validation studies: \"\n    f\"{len(fold2_validation_ids)}\"\n)\n\noverlap = (\n    fold2_training_ids\n    .intersection(fold2_validation_ids)\n)\n\nprint(\n    f\"Train/validation overlap: \"\n    f\"{len(overlap)}\"\n)\n\nif len(fold2_training_ids) != 39:\n    raise RuntimeError(\n        f\"Expected 39 training studies, \"\n        f\"found {len(fold2_training_ids)}.\"\n    )\n\nif len(fold2_validation_ids) != 19:\n    raise RuntimeError(\n        f\"Expected 19 validation studies, \"\n        f\"found {len(fold2_validation_ids)}.\"\n    )\n\nif overlap:\n    raise RuntimeError(\n        \"Train/validation study overlap detected.\"\n    )\n\nprint(\"Fold-2 split: PASS\")\n\n# ----------------------------------------------------------------\n# CREATE FOLD-2 TABLES\n# ----------------------------------------------------------------\n\nfold2_train_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].isin(\n        fold2_training_ids\n    )\n].copy()\n\nfold2_val_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].isin(\n        fold2_validation_ids\n    )\n].copy()\n\n# Preserve deterministic ordering\nfold2_train_table = (\n    fold2_train_table\n    .sort_values(\"StudyInstanceUID\")\n    .reset_index(drop=True)\n)\n\nfold2_val_table = (\n    fold2_val_table\n    .sort_values(\"StudyInstanceUID\")\n    .reset_index(drop=True)\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 TABLES\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Training table: \"\n    f\"{fold2_train_table.shape}\"\n)\n\nprint(\n    f\"Validation table: \"\n    f\"{fold2_val_table.shape}\"\n)\n\nif len(fold2_train_table) != 39:\n    raise RuntimeError(\n        \"Fold-2 training table must contain 39 studies.\"\n    )\n\nif len(fold2_val_table) != 19:\n    raise RuntimeError(\n        \"Fold-2 validation table must contain 19 studies.\"\n    )\n\n# ----------------------------------------------------------------\n# VERIFY TARGET VALUES\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TARGET VALIDATION\")\nprint(\"=\" * 70)\n\nfor target in TARGETS:\n\n    train_values = set(\n        fold2_train_table[target]\n        .astype(float)\n        .unique()\n    )\n\n    val_values = set(\n        fold2_val_table[target]\n        .astype(float)\n        .unique()\n    )\n\n    if not train_values.issubset({0.0, 1.0}):\n        raise RuntimeError(\n            f\"Invalid training values for {target}: \"\n            f\"{train_values}\"\n        )\n\n    if not val_values.issubset({0.0, 1.0}):\n        raise RuntimeError(\n            f\"Invalid validation values for {target}: \"\n            f\"{val_values}\"\n        )\n\n    print(\n        f\"{target:20s} \"\n        f\"train={sorted(train_values)} \"\n        f\"val={sorted(val_values)}\"\n    )\n\nprint(\"Target values: PASS\")\n\n# ----------------------------------------------------------------\n# CONSTRUCT DATASETS USING THE EXISTING CELL-33 CLASS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CREATING FOLD 2 DATASETS\")\nprint(\"=\" * 70)\n\ntrain_dataset = KneeStudyDataset(\n    fold2_train_table,\n    targets=TARGETS\n)\n\nval_dataset = KneeStudyDataset(\n    fold2_val_table,\n    targets=TARGETS\n)\n\nprint(\n    f\"Train Dataset: \"\n    f\"{len(train_dataset)}\"\n)\n\nprint(\n    f\"Validation Dataset: \"\n    f\"{len(val_dataset)}\"\n)\n\nif len(train_dataset) != 39:\n    raise RuntimeError(\n        \"Fold-2 training dataset length is not 39.\"\n    )\n\nif len(val_dataset) != 19:\n    raise RuntimeError(\n        \"Fold-2 validation dataset length is not 19.\"\n    )\n\n# ----------------------------------------------------------------\n# CREATE DATALOADERS\n# ----------------------------------------------------------------\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=0,\n    pin_memory=False\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False\n)\n\nprint(\n    f\"Train batches: \"\n    f\"{len(train_loader)}\"\n)\n\nprint(\n    f\"Validation batches: \"\n    f\"{len(val_loader)}\"\n)\n\n# ----------------------------------------------------------------\n# DATALOADER SMOKE TEST\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 DATALOADER SMOKE TEST\")\nprint(\"=\" * 70)\n\nbatch = next(\n    iter(train_loader)\n)\n\nif not isinstance(batch, dict):\n    raise RuntimeError(\n        f\"Expected dictionary batch, \"\n        f\"received {type(batch)}.\"\n    )\n\nprint(\n    \"Batch keys:\",\n    list(batch.keys())\n)\n\nif \"image\" not in batch:\n    raise RuntimeError(\n        \"Batch is missing image field.\"\n    )\n\nif \"label\" not in batch:\n    raise RuntimeError(\n        \"Batch is missing label field.\"\n    )\n\nimages = batch[\"image\"]\nlabels = batch[\"label\"]\n\nprint(\n    f\"Image batch shape: \"\n    f\"{tuple(images.shape)}\"\n)\n\nprint(\n    f\"Label batch shape: \"\n    f\"{tuple(labels.shape)}\"\n)\n\nexpected_image_shape = (\n    BATCH_SIZE,\n    EXPECTED_CHANNELS,\n    IMAGE_SIZE,\n    IMAGE_SIZE\n)\n\nexpected_label_shape = (\n    BATCH_SIZE,\n    len(TARGETS)\n)\n\nif tuple(images.shape) != expected_image_shape:\n    raise RuntimeError(\n        f\"Unexpected image shape. \"\n        f\"Expected {expected_image_shape}, \"\n        f\"got {tuple(images.shape)}.\"\n    )\n\nif tuple(labels.shape) != expected_label_shape:\n    raise RuntimeError(\n        f\"Unexpected label shape. \"\n        f\"Expected {expected_label_shape}, \"\n        f\"got {tuple(labels.shape)}.\"\n    )\n\nif images.dtype != torch.float32:\n    raise RuntimeError(\n        f\"Expected float32 images, \"\n        f\"got {images.dtype}.\"\n    )\n\nif labels.dtype != torch.float32:\n    raise RuntimeError(\n        f\"Expected float32 labels, \"\n        f\"got {labels.dtype}.\"\n    )\n\nif not torch.isfinite(images).all():\n    raise RuntimeError(\n        \"Image batch contains NaN or Inf.\"\n    )\n\nif not torch.isfinite(labels).all():\n    raise RuntimeError(\n        \"Label batch contains NaN or Inf.\"\n    )\n\n# ----------------------------------------------------------------\n# VALUE RANGE CHECK\n# ----------------------------------------------------------------\n\nprint(\n    f\"Image min: \"\n    f\"{images.min().item():.6f}\"\n)\n\nprint(\n    f\"Image max: \"\n    f\"{images.max().item():.6f}\"\n)\n\nprint(\n    f\"Image mean: \"\n    f\"{images.mean().item():.6f}\"\n)\n\nprint(\n    f\"Image std: \"\n    f\"{images.std().item():.6f}\"\n)\n\nif images.min().item() < -1e-6:\n    raise RuntimeError(\n        \"Image values below expected [0, 1] range.\"\n    )\n\nif images.max().item() > 1.000001:\n    raise RuntimeError(\n        \"Image values above expected [0, 1] range.\"\n    )\n\n# ----------------------------------------------------------------\n# STUDY IDS\n# ----------------------------------------------------------------\n\nif \"study_id\" in batch:\n    print(\n        \"First study IDs:\",\n        batch[\"study_id\"]\n    )\n\n# ----------------------------------------------------------------\n# MEMORY ESTIMATE\n# ----------------------------------------------------------------\n\nbytes_per_study = (\n    EXPECTED_CHANNELS\n    * IMAGE_SIZE\n    * IMAGE_SIZE\n    * 4\n)\n\nmb_per_study = (\n    bytes_per_study / (1024 ** 2)\n)\n\nmb_per_batch = (\n    mb_per_study * BATCH_SIZE\n)\n\nprint(\n    f\"Approximate image memory per study: \"\n    f\"{mb_per_study:.2f} MB\"\n)\n\nprint(\n    f\"Approximate image memory per batch: \"\n    f\"{mb_per_batch:.2f} MB\"\n)\n\n# ----------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 43 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"Fold-2 study split: PASS\")\nprint(\"Fold-2 datasets: PASS\")\nprint(\"Fold-2 DataLoaders: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target labels: PASS\")\nprint(\"Image dtype/range: PASS\")\nprint(\"NaN/Inf check: PASS\")\nprint(\"Lazy loading preserved: PASS\")\nprint(\"No model training performed.\")\nprint(\"Ready for Fold-2 model reset and training.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:13:35.881115Z","iopub.execute_input":"2026-08-11T10:13:35.881438Z","iopub.status.idle":"2026-08-11T10:13:36.714517Z","shell.execute_reply.started":"2026-08-11T10:13:35.881412Z","shell.execute_reply":"2026-08-11T10:13:36.713196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 44 - FOLD 2 TRAINING\n# ================================================================\n\nimport os\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.metrics import f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 44 - FOLD 2 TRAINING\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nSEED = 42\nFOLD = 2\n\nMAX_EPOCHS = 10\nEARLY_STOPPING_PATIENCE = 3\n\nLEARNING_RATE = 0.001\nWEIGHT_DECAY = 0.0001\nMAX_GRAD_NORM = 1.0\n\n# CPU, because previous cells confirmed CUDA is unavailable\nDEVICE = torch.device(\"cpu\")\n\nCHECKPOINT_PATH = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell44_fold2_best_model.pt\"\n)\n\nHISTORY_PATH = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell44_fold2_history.csv\"\n)\n\nMETRICS_PATH = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell44_fold2_metrics.csv\"\n)\n\n# ----------------------------------------------------------------\n# RANDOM SEEDS\n# ----------------------------------------------------------------\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nprint(f\"Fold: {FOLD}\")\nprint(f\"Maximum epochs: {MAX_EPOCHS}\")\nprint(f\"Early stopping patience: {EARLY_STOPPING_PATIENCE}\")\nprint(f\"Device: {DEVICE}\")\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECT CHECK\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"train_loader\",\n    \"val_loader\",\n    \"TARGETS\",\n    \"fold2_train_table\",\n    \"fold2_val_table\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# VERIFY FOLD-2 DATA\n# ----------------------------------------------------------------\n\nif len(fold2_train_table) != 39:\n    raise RuntimeError(\n        f\"Expected 39 Fold-2 training studies, \"\n        f\"found {len(fold2_train_table)}.\"\n    )\n\nif len(fold2_val_table) != 19:\n    raise RuntimeError(\n        f\"Expected 19 Fold-2 validation studies, \"\n        f\"found {len(fold2_val_table)}.\"\n    )\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD DATA\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Training studies: {len(fold2_train_table)}\"\n)\n\nprint(\n    f\"Validation studies: {len(fold2_val_table)}\"\n)\n\nprint(\n    f\"Training batches: {len(train_loader)}\"\n)\n\nprint(\n    f\"Validation batches: {len(val_loader)}\"\n)\n\n# ----------------------------------------------------------------\n# RESET TRAINABLE MODEL PARAMETERS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"RESETTING MODEL FOR FOLD 2\")\nprint(\"=\" * 70)\n\n# The Cell-35 architecture has a frozen ResNet18 backbone\n# and trainable final classification layer.\n#\n# We reset ONLY parameters that were trainable before training.\n# Frozen pretrained backbone weights are preserved.\n\ntrainable_parameter_count = 0\nreset_parameter_count = 0\nreset_tensor_count = 0\n\nfor module in model.modules():\n\n    # Reset standard trainable layers where a reset_parameters\n    # method exists.\n    if hasattr(module, \"reset_parameters\"):\n\n        module_trainable = any(\n            parameter.requires_grad\n            for parameter in module.parameters(recurse=False)\n        )\n\n        if module_trainable:\n            before = sum(\n                parameter.numel()\n                for parameter in module.parameters(recurse=False)\n                if parameter.requires_grad\n            )\n\n            if before > 0:\n                module.reset_parameters()\n\n                reset_tensor_count += 1\n                reset_parameter_count += before\n\n# Count trainable parameters after reset\nfor parameter in model.parameters():\n    if parameter.requires_grad:\n        trainable_parameter_count += parameter.numel()\n\nprint(\n    f\"Trainable parameters after reset: \"\n    f\"{trainable_parameter_count:,}\"\n)\n\nprint(\n    f\"Parameters reset: \"\n    f\"{reset_parameter_count:,}\"\n)\n\nprint(\n    f\"Trainable modules reset: \"\n    f\"{reset_tensor_count}\"\n)\n\nif trainable_parameter_count != 200268:\n    raise RuntimeError(\n        \"Unexpected trainable parameter count. \"\n        f\"Expected 200,268, found \"\n        f\"{trainable_parameter_count}.\"\n    )\n\nif reset_parameter_count != 200268:\n    raise RuntimeError(\n        \"Unexpected reset parameter count. \"\n        f\"Expected 200,268, found \"\n        f\"{reset_parameter_count}.\"\n    )\n\nprint(\"Trainable parameter count: PASS\")\nprint(\"Trainable parameter reset: PASS\")\n\n# ----------------------------------------------------------------\n# MOVE MODEL TO CPU\n# ----------------------------------------------------------------\n\nmodel = model.to(DEVICE)\n\n# ----------------------------------------------------------------\n# COMPUTE FOLD-2 CLASS WEIGHTS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 CLASS WEIGHTS\")\nprint(\"=\" * 70)\n\npositive_counts = []\nnegative_counts = []\npositive_weights = []\n\nfor target in TARGETS:\n\n    positives = int(\n        fold2_train_table[target]\n        .astype(float)\n        .sum()\n    )\n\n    negatives = int(\n        len(fold2_train_table) - positives\n    )\n\n    if positives <= 0:\n        raise RuntimeError(\n            f\"No positive samples for target: {target}\"\n        )\n\n    if negatives <= 0:\n        raise RuntimeError(\n            f\"No negative samples for target: {target}\"\n        )\n\n    pos_weight = (\n        negatives / positives\n    )\n\n    positive_counts.append(positives)\n    negative_counts.append(negatives)\n    positive_weights.append(pos_weight)\n\n    print(\n        f\"{target:20s} \"\n        f\"positive={positives:2d} \"\n        f\"negative={negatives:2d} \"\n        f\"pos_weight={pos_weight:.4f}\"\n    )\n\npos_weight_tensor = torch.tensor(\n    positive_weights,\n    dtype=torch.float32,\n    device=DEVICE\n)\n\n# ----------------------------------------------------------------\n# LOSS\n# ----------------------------------------------------------------\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_tensor\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINING CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(\"Loss: BCEWithLogitsLoss\")\nprint(\"Optimizer: AdamW\")\nprint(f\"Learning rate: {LEARNING_RATE}\")\nprint(f\"Weight decay: {WEIGHT_DECAY}\")\nprint(f\"Gradient clipping: {MAX_GRAD_NORM}\")\nprint(\"Scheduler: ReduceLROnPlateau\")\n\n# ----------------------------------------------------------------\n# OPTIMIZER\n# ----------------------------------------------------------------\n\ntrainable_parameters = [\n    parameter\n    for parameter in model.parameters()\n    if parameter.requires_grad\n]\n\noptimizer = AdamW(\n    trainable_parameters,\n    lr=LEARNING_RATE,\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler = ReduceLROnPlateau(\n    optimizer,\n    mode=\"min\",\n    factor=0.5,\n    patience=2\n)\n\n# ----------------------------------------------------------------\n# TRAINING HISTORY\n# ----------------------------------------------------------------\n\nhistory = []\n\nbest_val_loss = float(\"inf\")\nbest_epoch = 0\n\nepochs_without_improvement = 0\n\n# ----------------------------------------------------------------\n# TRAINING LOOP\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING FOLD 2 TRAINING\")\nprint(\"=\" * 70)\n\nfor epoch in range(1, MAX_EPOCHS + 1):\n\n    model.train()\n\n    train_losses = []\n\n    for batch in train_loader:\n\n        images = batch[\"image\"].to(\n            DEVICE,\n            non_blocking=False\n        )\n\n        labels = batch[\"label\"].to(\n            DEVICE,\n            non_blocking=False\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        logits = model(images)\n\n        if logits.shape != labels.shape:\n            raise RuntimeError(\n                \"Logit/label shape mismatch: \"\n                f\"{tuple(logits.shape)} vs \"\n                f\"{tuple(labels.shape)}\"\n            )\n\n        loss = criterion(\n            logits,\n            labels\n        )\n\n        if not torch.isfinite(loss):\n            raise RuntimeError(\n                \"Non-finite training loss detected.\"\n            )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            trainable_parameters,\n            MAX_GRAD_NORM\n        )\n\n        optimizer.step()\n\n        train_losses.append(\n            loss.detach().item()\n        )\n\n    # ------------------------------------------------------------\n    # VALIDATION\n    # ------------------------------------------------------------\n\n    model.eval()\n\n    val_losses = []\n\n    all_probabilities = []\n    all_labels = []\n\n    with torch.no_grad():\n\n        for batch in val_loader:\n\n            images = batch[\"image\"].to(\n                DEVICE,\n                non_blocking=False\n            )\n\n            labels = batch[\"label\"].to(\n                DEVICE,\n                non_blocking=False\n            )\n\n            logits = model(images)\n\n            loss = criterion(\n                logits,\n                labels\n            )\n\n            probabilities = torch.sigmoid(\n                logits\n            )\n\n            if not torch.isfinite(loss):\n                raise RuntimeError(\n                    \"Non-finite validation loss detected.\"\n                )\n\n            if not torch.isfinite(probabilities).all():\n                raise RuntimeError(\n                    \"Non-finite validation probabilities detected.\"\n                )\n\n            val_losses.append(\n                loss.item()\n            )\n\n            all_probabilities.append(\n                probabilities.cpu().numpy()\n            )\n\n            all_labels.append(\n                labels.cpu().numpy()\n            )\n\n    train_loss = float(\n        np.mean(train_losses)\n    )\n\n    val_loss = float(\n        np.mean(val_losses)\n    )\n\n    probabilities_np = np.concatenate(\n        all_probabilities,\n        axis=0\n    )\n\n    labels_np = np.concatenate(\n        all_labels,\n        axis=0\n    )\n\n    predictions_np = (\n        probabilities_np >= 0.50\n    ).astype(int)\n\n    target_f1 = []\n\n    for target_index in range(len(TARGETS)):\n\n        target_f1.append(\n            f1_score(\n                labels_np[:, target_index],\n                predictions_np[:, target_index],\n                zero_division=0\n            )\n        )\n\n    mean_f1 = float(\n        np.mean(target_f1)\n    )\n\n    # ------------------------------------------------------------\n    # SCHEDULER\n    # ------------------------------------------------------------\n\n    scheduler.step(val_loss)\n\n    current_lr = optimizer.param_groups[0][\"lr\"]\n\n    # ------------------------------------------------------------\n    # HISTORY\n    # ------------------------------------------------------------\n\n    history.append(\n        {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"mean_f1\": mean_f1,\n            \"learning_rate\": current_lr\n        }\n    )\n\n    # ------------------------------------------------------------\n    # BEST CHECKPOINT\n    # ------------------------------------------------------------\n\n    if val_loss < best_val_loss:\n\n        best_val_loss = val_loss\n        best_epoch = epoch\n        epochs_without_improvement = 0\n\n        torch.save(\n            {\n                \"model_state_dict\": model.state_dict(),\n                \"epoch\": epoch,\n                \"val_loss\": val_loss,\n                \"targets\": TARGETS,\n                \"fold\": FOLD,\n                \"seed\": SEED\n            },\n            CHECKPOINT_PATH\n        )\n\n    else:\n\n        epochs_without_improvement += 1\n\n    print(f\"\\nEpoch {epoch:02d}/{MAX_EPOCHS}\")\n    print(f\"Train Loss: {train_loss:.5f}\")\n    print(f\"Val Loss:   {val_loss:.5f}\")\n    print(f\"Mean F1:    {mean_f1:.4f}\")\n    print(f\"Learning Rate: {current_lr:.6f}\")\n    print(f\"Best Epoch: {best_epoch}\")\n    print(f\"Best Val Loss: {best_val_loss:.5f}\")\n\n    # ------------------------------------------------------------\n    # EARLY STOPPING\n    # ------------------------------------------------------------\n\n    if (\n        epochs_without_improvement\n        >= EARLY_STOPPING_PATIENCE\n    ):\n\n        print(\n            \"\\nEarly stopping triggered \"\n            f\"after {EARLY_STOPPING_PATIENCE} \"\n            \"epochs without improvement.\"\n        )\n\n        break\n\n# ----------------------------------------------------------------\n# RESTORE BEST CHECKPOINT\n# ----------------------------------------------------------------\n\nif not os.path.exists(CHECKPOINT_PATH):\n    raise RuntimeError(\n        \"Best Fold-2 checkpoint was not created.\"\n    )\n\ncheckpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\n# ----------------------------------------------------------------\n# FINAL BEST-CHECKPOINT VALIDATION\n# ----------------------------------------------------------------\n\nfinal_val_losses = []\nfinal_probabilities = []\nfinal_labels = []\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        images = batch[\"image\"].to(\n            DEVICE\n        )\n\n        labels = batch[\"label\"].to(\n            DEVICE\n        )\n\n        logits = model(images)\n\n        loss = criterion(\n            logits,\n            labels\n        )\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n        final_val_losses.append(\n            loss.item()\n        )\n\n        final_probabilities.append(\n            probabilities.cpu().numpy()\n        )\n\n        final_labels.append(\n            labels.cpu().numpy()\n        )\n\nfinal_val_loss = float(\n    np.mean(final_val_losses)\n)\n\nfinal_probabilities = np.concatenate(\n    final_probabilities,\n    axis=0\n)\n\nfinal_labels = np.concatenate(\n    final_labels,\n    axis=0\n)\n\nfinal_predictions = (\n    final_probabilities >= 0.50\n).astype(int)\n\n# ----------------------------------------------------------------\n# PER-TARGET METRICS\n# ----------------------------------------------------------------\n\nmetrics = []\n\nfor target_index, target in enumerate(TARGETS):\n\n    y_true = final_labels[:, target_index]\n    y_pred = final_predictions[:, target_index]\n\n    tp = int(\n        np.sum(\n            (y_true == 1) &\n            (y_pred == 1)\n        )\n    )\n\n    tn = int(\n        np.sum(\n            (y_true == 0) &\n            (y_pred == 0)\n        )\n    )\n\n    fp = int(\n        np.sum(\n            (y_true == 0) &\n            (y_pred == 1)\n        )\n    )\n\n    fn = int(\n        np.sum(\n            (y_true == 1) &\n            (y_pred == 0)\n        )\n    )\n\n    actual_positive = tp + fn\n    actual_negative = tn + fp\n\n    sensitivity = (\n        tp / actual_positive\n        if actual_positive > 0\n        else 0.0\n    )\n\n    specificity = (\n        tn / actual_negative\n        if actual_negative > 0\n        else 0.0\n    )\n\n    precision = (\n        tp / (tp + fp)\n        if (tp + fp) > 0\n        else 0.0\n    )\n\n    f1 = (\n        2 * precision * sensitivity\n        / (precision + sensitivity)\n        if (precision + sensitivity) > 0\n        else 0.0\n    )\n\n    metrics.append(\n        {\n            \"target\": target,\n            \"positive_count\": int(np.sum(y_true)),\n            \"negative_count\": int(np.sum(y_true == 0)),\n            \"predicted_positive\": int(np.sum(y_pred)),\n            \"TP\": tp,\n            \"TN\": tn,\n            \"FP\": fp,\n            \"FN\": fn,\n            \"sensitivity\": sensitivity,\n            \"specificity\": specificity,\n            \"precision\": precision,\n            \"f1\": f1,\n            \"mean_probability\": float(\n                np.mean(\n                    final_probabilities[:, target_index]\n                )\n            )\n        }\n    )\n\nmetrics_df = pd.DataFrame(metrics)\n\n# ----------------------------------------------------------------\n# SAVE RESULTS\n# ----------------------------------------------------------------\n\nhistory_df = pd.DataFrame(history)\n\nhistory_df.to_csv(\n    HISTORY_PATH,\n    index=False\n)\n\nmetrics_df.to_csv(\n    METRICS_PATH,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL SUMMARY\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 2 TRAINING COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Best epoch: {best_epoch}\"\n)\n\nprint(\n    f\"Best validation loss: \"\n    f\"{best_val_loss:.6f}\"\n)\n\nprint(\n    f\"Final validation loss: \"\n    f\"{final_val_loss:.6f}\"\n)\n\nprint(\"\\nPer-target validation metrics:\")\nprint(\n    metrics_df[\n        [\n            \"target\",\n            \"positive_count\",\n            \"negative_count\",\n            \"predicted_positive\",\n            \"sensitivity\",\n            \"specificity\",\n            \"precision\",\n            \"f1\"\n        ]\n    ].to_string(index=False)\n)\n\nprint(\n    f\"\\nMean F1: \"\n    f\"{metrics_df['f1'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean sensitivity: \"\n    f\"{metrics_df['sensitivity'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean specificity: \"\n    f\"{metrics_df['specificity'].mean():.4f}\"\n)\n\nprint(\n    f\"Mean precision: \"\n    f\"{metrics_df['precision'].mean():.4f}\"\n)\n\nprint(\n    f\"\\nCheckpoint saved: \"\n    f\"{CHECKPOINT_PATH}\"\n)\n\nprint(\n    f\"History saved: \"\n    f\"{HISTORY_PATH}\"\n)\n\nprint(\n    f\"Metrics saved: \"\n    f\"{METRICS_PATH}\"\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 44 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:14:56.53014Z","iopub.execute_input":"2026-08-11T10:14:56.530522Z","iopub.status.idle":"2026-08-11T10:18:43.637763Z","shell.execute_reply.started":"2026-08-11T10:14:56.530488Z","shell.execute_reply":"2026-08-11T10:18:43.636776Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 45 - FOLD 2 PROBABILITY + THRESHOLD AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score, f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 45 - FOLD 2 PROBABILITY + THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"val_loader\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\nDEVICE = torch.device(\"cpu\")\n\n# ----------------------------------------------------------------\n# LOAD BEST FOLD-2 CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell44_fold2_best_model.pt\"\n)\n\nif not os.path.exists(checkpoint_path):\n    raise RuntimeError(\n        \"Missing Fold-2 checkpoint: \"\n        + checkpoint_path\n    )\n\nprint(\"\\nLoading best Fold-2 checkpoint:\")\nprint(checkpoint_path)\n\ncheckpoint = torch.load(\n    checkpoint_path,\n    map_location=DEVICE\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel = model.to(DEVICE)\nmodel.eval()\n\nprint(\"Best Fold-2 checkpoint: LOADED\")\nprint(\"Model mode: evaluation\")\n\n# ----------------------------------------------------------------\n# COLLECT VALIDATION PREDICTIONS\n# ----------------------------------------------------------------\n\nall_probabilities = []\nall_labels = []\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        images = batch[\"image\"].to(DEVICE)\n        labels = batch[\"label\"].to(DEVICE)\n\n        logits = model(images)\n\n        probabilities = torch.sigmoid(logits)\n\n        if not torch.isfinite(probabilities).all():\n            raise RuntimeError(\n                \"NaN/Inf detected in probabilities.\"\n            )\n\n        all_probabilities.append(\n            probabilities.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.cpu().numpy()\n        )\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\nlabels = np.concatenate(\n    all_labels,\n    axis=0\n)\n\nprint(\n    f\"\\nProbability shape: {probabilities.shape}\"\n)\n\nprint(\n    f\"Label shape: {labels.shape}\"\n)\n\nif probabilities.shape != labels.shape:\n    raise RuntimeError(\n        \"Probability and label shapes do not match.\"\n    )\n\nif probabilities.shape[1] != len(TARGETS):\n    raise RuntimeError(\n        \"Unexpected number of target columns.\"\n    )\n\nprint(\"Probability validity: PASS\")\nprint(\"Shape consistency: PASS\")\n\n# ----------------------------------------------------------------\n# PROBABILITY SEPARATION + ROC-AUC\n# ----------------------------------------------------------------\n\nseparation_rows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    positive_probs = y_prob[y_true == 1]\n    negative_probs = y_prob[y_true == 0]\n\n    positive_mean = float(\n        np.mean(positive_probs)\n    )\n\n    negative_mean = float(\n        np.mean(negative_probs)\n    )\n\n    positive_median = float(\n        np.median(positive_probs)\n    )\n\n    negative_median = float(\n        np.median(negative_probs)\n    )\n\n    separation = (\n        positive_mean -\n        negative_mean\n    )\n\n    if (\n        len(np.unique(y_true)) == 2\n    ):\n        auc = float(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n    else:\n        auc = np.nan\n\n    separation_rows.append(\n        {\n            \"target\": target,\n            \"positive_count\": int(np.sum(y_true == 1)),\n            \"negative_count\": int(np.sum(y_true == 0)),\n            \"positive_mean\": positive_mean,\n            \"negative_mean\": negative_mean,\n            \"positive_median\": positive_median,\n            \"negative_median\": negative_median,\n            \"probability_separation\": separation,\n            \"roc_auc\": auc\n        }\n    )\n\nseparation_df = pd.DataFrame(\n    separation_rows\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PROBABILITY SEPARATION\")\nprint(\"=\" * 70)\n\nprint(\n    separation_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# THRESHOLD AUDIT\n# ----------------------------------------------------------------\n\nthresholds = np.arange(\n    0.20,\n    0.651,\n    0.025\n)\n\nthreshold_rows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    f1_at_050 = f1_score(\n        y_true,\n        (y_prob >= 0.50).astype(int),\n        zero_division=0\n    )\n\n    best_threshold = 0.50\n    best_f1 = f1_at_050\n\n    for threshold in thresholds:\n\n        y_pred = (\n            y_prob >= threshold\n        ).astype(int)\n\n        current_f1 = f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n\n        if current_f1 > best_f1:\n\n            best_f1 = current_f1\n            best_threshold = float(\n                threshold\n            )\n\n    threshold_rows.append(\n        {\n            \"target\": target,\n            \"f1_at_0.50\": float(f1_at_050),\n            \"best_threshold\": float(best_threshold),\n            \"best_f1\": float(best_f1),\n            \"f1_improvement\": float(\n                best_f1 - f1_at_050\n            )\n        }\n    )\n\nthreshold_df = pd.DataFrame(\n    threshold_rows\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"THRESHOLD AUDIT\")\nprint(\"=\" * 70)\n\nprint(\n    threshold_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------\n\nmean_auc = float(\n    separation_df[\"roc_auc\"].mean()\n)\n\nmean_separation = float(\n    separation_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    threshold_df[\"f1_at_0.50\"].mean()\n)\n\nmean_best_f1 = float(\n    threshold_df[\"best_f1\"].mean()\n)\n\nf1_improvement = (\n    mean_best_f1 -\n    mean_f1_050\n)\n\nauc_good_targets = separation_df[\n    separation_df[\"roc_auc\"] >= 0.65\n][\"target\"].tolist()\n\nnegative_separation_targets = separation_df[\n    separation_df[\"probability_separation\"] < 0\n][\"target\"].tolist()\n\nimprovement_targets = threshold_df[\n    threshold_df[\"f1_improvement\"] > 0.05\n][\"target\"].tolist()\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Mean ROC-AUC: \"\n    f\"{mean_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50: \"\n    f\"{mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{f1_improvement:.4f}\"\n)\n\nprint(\n    f\"\\nTargets with ROC-AUC >= 0.65: \"\n    f\"{len(auc_good_targets)}\"\n)\n\nprint(\n    auc_good_targets\n)\n\nprint(\n    f\"\\nTargets with negative probability separation: \"\n    f\"{len(negative_separation_targets)}\"\n)\n\nprint(\n    negative_separation_targets\n)\n\nprint(\n    f\"\\nTargets with >0.05 F1 improvement: \"\n    f\"{len(improvement_targets)}\"\n)\n\nprint(\n    improvement_targets\n)\n\n# ----------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------\n\nseparation_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell45_fold2_probability_separation.csv\"\n)\n\nthreshold_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell45_fold2_threshold_audit.csv\"\n)\n\nseparation_df.to_csv(\n    separation_path,\n    index=False\n)\n\nthreshold_df.to_csv(\n    threshold_path,\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(separation_path)\nprint(threshold_path)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 45 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:19:38.270991Z","iopub.execute_input":"2026-08-11T10:19:38.271353Z","iopub.status.idle":"2026-08-11T10:19:47.982643Z","shell.execute_reply.started":"2026-08-11T10:19:38.271322Z","shell.execute_reply":"2026-08-11T10:19:47.981618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 46 - 3-FOLD OUT-OF-FOLD PREDICTION CONSOLIDATION\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score, f1_score\n\nprint(\"=\" * 70)\nprint(\"CELL 46 - 3-FOLD OUT-OF-FOLD PREDICTION CONSOLIDATION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONSTANTS\n# ----------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nMODELING_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_modeling_studies.csv\"\n)\n\nFOLD_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell29_study_folds.csv\"\n)\n\nCHECKPOINTS = {\n    0: os.path.join(\n        AUDIT_DIR,\n        \"cell37_fold0_best_model.pt\"\n    ),\n    1: os.path.join(\n        AUDIT_DIR,\n        \"cell40_fold1_best_model.pt\"\n    ),\n    2: os.path.join(\n        AUDIT_DIR,\n        \"cell44_fold2_best_model.pt\"\n    )\n}\n\nTARGETS_EXPECTED = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\nDEVICE = torch.device(\"cpu\")\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"KneeStudyDataset\",\n    \"model\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# LOAD MODELING TABLE AND FOLD TABLE FROM SAVED FILES\n# ----------------------------------------------------------------\n\nif not os.path.exists(MODELING_PATH):\n    raise RuntimeError(\n        \"Missing modeling table: \"\n        + MODELING_PATH\n    )\n\nif not os.path.exists(FOLD_PATH):\n    raise RuntimeError(\n        \"Missing fold table: \"\n        + FOLD_PATH\n    )\n\nmodeling_table = pd.read_csv(\n    MODELING_PATH\n)\n\nstudy_folds = pd.read_csv(\n    FOLD_PATH\n)\n\nprint(\"\\nModeling table:\")\nprint(modeling_table.shape)\n\nprint(\"Fold table:\")\nprint(study_folds.shape)\n\n# ----------------------------------------------------------------\n# BASIC TABLE VALIDATION\n# ----------------------------------------------------------------\n\nif \"StudyInstanceUID\" not in modeling_table.columns:\n    raise RuntimeError(\n        \"StudyInstanceUID missing from modeling table.\"\n    )\n\nif \"StudyInstanceUID\" not in study_folds.columns:\n    raise RuntimeError(\n        \"StudyInstanceUID missing from fold table.\"\n    )\n\nif \"fold\" not in study_folds.columns:\n    raise RuntimeError(\n        \"fold column missing from fold table.\"\n    )\n\nmissing_targets = [\n    target\n    for target in TARGETS_EXPECTED\n    if target not in modeling_table.columns\n]\n\nif missing_targets:\n    raise RuntimeError(\n        \"Missing target columns: \"\n        + \", \".join(missing_targets)\n    )\n\n# ----------------------------------------------------------------\n# PRIMARY SERIES COLUMN VALIDATION\n# ----------------------------------------------------------------\n\nprimary_columns = [\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\"\n]\n\nmissing_primary_columns = [\n    column\n    for column in primary_columns\n    if column not in modeling_table.columns\n]\n\nif missing_primary_columns:\n    raise RuntimeError(\n        \"Missing primary-series columns: \"\n        + \", \".join(missing_primary_columns)\n    )\n\nprint(\"Target columns: PASS\")\nprint(\"Primary-series columns: PASS\")\n\n# ----------------------------------------------------------------\n# MERGE FOLD ASSIGNMENTS\n# ----------------------------------------------------------------\n\nfold_columns = [\n    \"StudyInstanceUID\",\n    \"fold\"\n]\n\nfold_assignment = study_folds[\n    fold_columns\n].copy()\n\nif len(\n    fold_assignment[\"StudyInstanceUID\"].unique()\n) != len(fold_assignment):\n\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID values \"\n        \"found in fold table.\"\n    )\n\nmodeling_table = modeling_table.drop(\n    columns=[\"fold\"],\n    errors=\"ignore\"\n)\n\nmodeling_table = modeling_table.merge(\n    fold_assignment,\n    on=\"StudyInstanceUID\",\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\nif modeling_table[\"fold\"].isna().any():\n    raise RuntimeError(\n        \"Some modeling studies do not have \"\n        \"a fold assignment.\"\n    )\n\nprint(\"Fold assignment merge: PASS\")\n\n# ----------------------------------------------------------------\n# GLOBAL STUDY VALIDATION\n# ----------------------------------------------------------------\n\nif len(modeling_table) != 58:\n    raise RuntimeError(\n        f\"Expected 58 modeling studies, \"\n        f\"found {len(modeling_table)}.\"\n    )\n\nif (\n    modeling_table[\"StudyInstanceUID\"]\n    .nunique()\n    != 58\n):\n    raise RuntimeError(\n        \"Modeling table does not contain \"\n        \"58 unique studies.\"\n    )\n\navailable_folds = sorted(\n    modeling_table[\"fold\"].astype(int).unique()\n)\n\nif available_folds != [0, 1, 2]:\n    raise RuntimeError(\n        \"Expected folds [0, 1, 2], \"\n        f\"found {available_folds}\"\n    )\n\nprint(\"Total modeling studies: 58\")\nprint(\"Unique studies: 58\")\nprint(\"Available folds: [0, 1, 2]\")\nprint(\"Three-fold configuration: PASS\")\n\n# ----------------------------------------------------------------\n# CHECK CHECKPOINTS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CHECKPOINT VALIDATION\")\nprint(\"=\" * 70)\n\nfor fold, path in CHECKPOINTS.items():\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Missing Fold-{fold} checkpoint: {path}\"\n        )\n\n    print(\n        f\"Fold {fold} checkpoint: EXISTS\"\n    )\n\nprint(\"All three checkpoints: PASS\")\n\n# ----------------------------------------------------------------\n# FUNCTION TO CREATE VALIDATION LOADER\n# ----------------------------------------------------------------\n\ndef create_validation_loader(\n    fold_number,\n    dataframe\n):\n\n    validation_table = dataframe[\n        dataframe[\"fold\"].astype(int)\n        == fold_number\n    ].copy()\n\n    if validation_table.empty:\n        raise RuntimeError(\n            f\"Fold {fold_number} validation table is empty.\"\n        )\n\n    dataset = KneeStudyDataset(\n        validation_table,\n        targets=TARGETS_EXPECTED\n    )\n\n    loader = torch.utils.data.DataLoader(\n        dataset,\n        batch_size=2,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=False\n    )\n\n    return validation_table, dataset, loader\n\n# ----------------------------------------------------------------\n# OOF STORAGE\n# ----------------------------------------------------------------\n\noof_rows = []\n\n# ----------------------------------------------------------------\n# PROCESS EACH FOLD\n# ----------------------------------------------------------------\n\nfor fold_number in [0, 1, 2]:\n\n    print(\"\\n\" + \"=\" * 70)\n    print(\n        f\"PROCESSING FOLD {fold_number} VALIDATION\"\n    )\n    print(\"=\" * 70)\n\n    validation_table, validation_dataset, validation_loader = (\n        create_validation_loader(\n            fold_number,\n            modeling_table\n        )\n    )\n\n    expected_studies = (\n        validation_table[\n            \"StudyInstanceUID\"\n        ]\n        .astype(str)\n        .tolist()\n    )\n\n    print(\n        f\"Validation studies: \"\n        f\"{len(expected_studies)}\"\n    )\n\n    print(\n        f\"Validation dataset length: \"\n        f\"{len(validation_dataset)}\"\n    )\n\n    # ------------------------------------------------------------\n    # LOAD CHECKPOINT\n    # ------------------------------------------------------------\n\n    checkpoint_path = CHECKPOINTS[\n        fold_number\n    ]\n\n    checkpoint = torch.load(\n        checkpoint_path,\n        map_location=DEVICE\n    )\n\n    model.load_state_dict(\n        checkpoint[\"model_state_dict\"]\n    )\n\n    model = model.to(DEVICE)\n    model.eval()\n\n    print(\n        f\"Fold-{fold_number} checkpoint: LOADED\"\n    )\n\n    # ------------------------------------------------------------\n    # PREDICTION\n    # ------------------------------------------------------------\n\n    fold_study_ids = []\n\n    with torch.no_grad():\n\n        for batch in validation_loader:\n\n            images = batch[\"image\"].to(DEVICE)\n\n            labels = batch[\"label\"].cpu().numpy()\n\n            study_ids = [\n                str(study_id)\n                for study_id in batch[\"study_id\"]\n            ]\n\n            logits = model(images)\n\n            probabilities = torch.sigmoid(\n                logits\n            ).cpu().numpy()\n\n            if not np.isfinite(\n                probabilities\n            ).all():\n\n                raise RuntimeError(\n                    f\"NaN/Inf detected in \"\n                    f\"Fold {fold_number} predictions.\"\n                )\n\n            if probabilities.shape[1] != 12:\n                raise RuntimeError(\n                    f\"Unexpected output shape \"\n                    f\"in Fold {fold_number}: \"\n                    f\"{probabilities.shape}\"\n                )\n\n            for row_index, study_id in enumerate(\n                study_ids\n            ):\n\n                fold_study_ids.append(\n                    study_id\n                )\n\n                row = {\n                    \"StudyInstanceUID\": study_id,\n                    \"fold\": fold_number\n                }\n\n                for target_index, target in enumerate(\n                    TARGETS_EXPECTED\n                ):\n\n                    row[\n                        f\"{target}_label\"\n                    ] = float(\n                        labels[\n                            row_index,\n                            target_index\n                        ]\n                    )\n\n                    row[\n                        f\"{target}_prob\"\n                    ] = float(\n                        probabilities[\n                            row_index,\n                            target_index\n                        ]\n                    )\n\n                oof_rows.append(row)\n\n    # ------------------------------------------------------------\n    # FOLD COVERAGE VALIDATION\n    # ------------------------------------------------------------\n\n    if len(fold_study_ids) != len(\n        expected_studies\n    ):\n\n        raise RuntimeError(\n            f\"Fold {fold_number}: expected \"\n            f\"{len(expected_studies)} predictions, \"\n            f\"got {len(fold_study_ids)}.\"\n        )\n\n    if len(set(fold_study_ids)) != len(\n        fold_study_ids\n    ):\n\n        raise RuntimeError(\n            f\"Fold {fold_number}: duplicate \"\n            f\"study predictions detected.\"\n        )\n\n    if set(fold_study_ids) != set(\n        expected_studies\n    ):\n\n        raise RuntimeError(\n            f\"Fold {fold_number}: predicted study \"\n            f\"IDs do not match validation studies.\"\n        )\n\n    print(\n        f\"Fold {fold_number} predictions: \"\n        f\"{len(fold_study_ids)}\"\n    )\n\n    print(\n        f\"Fold {fold_number} coverage: PASS\"\n    )\n\n# ----------------------------------------------------------------\n# BUILD OOF DATAFRAME\n# ----------------------------------------------------------------\n\noof_df = pd.DataFrame(\n    oof_rows\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OOF DATASET VALIDATION\")\nprint(\"=\" * 70)\n\nprint(\n    \"OOF shape:\",\n    oof_df.shape\n)\n\nif len(oof_df) != 58:\n    raise RuntimeError(\n        f\"Expected 58 OOF rows, \"\n        f\"found {len(oof_df)}.\"\n    )\n\nif (\n    oof_df[\"StudyInstanceUID\"]\n    .nunique()\n    != 58\n):\n\n    raise RuntimeError(\n        \"OOF predictions do not contain \"\n        \"exactly 58 unique studies.\"\n    )\n\nif (\n    oof_df.groupby(\n        \"StudyInstanceUID\"\n    ).size() != 1\n).any():\n\n    raise RuntimeError(\n        \"Some studies appear more than once \"\n        \"in the OOF prediction table.\"\n    )\n\nprint(\"OOF rows: 58\")\nprint(\"Unique OOF studies: 58\")\nprint(\"One prediction per study: PASS\")\n\n# ----------------------------------------------------------------\n# FOLD COVERAGE\n# ----------------------------------------------------------------\n\nprint(\"\\nOOF fold distribution:\")\n\nprint(\n    oof_df[\"fold\"]\n    .value_counts()\n    .sort_index()\n)\n\nexpected_fold_sizes = {\n    0: 20,\n    1: 19,\n    2: 19\n}\n\nfor fold_number, expected_size in (\n    expected_fold_sizes.items()\n):\n\n    actual_size = int(\n        (oof_df[\"fold\"] == fold_number).sum()\n    )\n\n    if actual_size != expected_size:\n\n        raise RuntimeError(\n            f\"Fold {fold_number}: expected \"\n            f\"{expected_size} OOF rows, \"\n            f\"found {actual_size}.\"\n        )\n\nprint(\"Fold coverage: PASS\")\n\n# ----------------------------------------------------------------\n# CALCULATE OOF METRICS\n# ----------------------------------------------------------------\n\nmetric_rows = []\n\nfor target in TARGETS_EXPECTED:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].values\n\n    y_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    auc = roc_auc_score(\n        y_true,\n        y_prob\n    )\n\n    f1 = f1_score(\n        y_true,\n        y_pred,\n        zero_division=0\n    )\n\n    metric_rows.append(\n        {\n            \"target\": target,\n            \"positive_count\": int(\n                y_true.sum()\n            ),\n            \"negative_count\": int(\n                len(y_true) - y_true.sum()\n            ),\n            \"roc_auc\": float(auc),\n            \"f1_at_0.50\": float(f1)\n        }\n    )\n\noof_metrics = pd.DataFrame(\n    metric_rows\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OOF BASELINE METRICS\")\nprint(\"=\" * 70)\n\nprint(\n    oof_metrics.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\nmean_oof_auc = float(\n    oof_metrics[\"roc_auc\"].mean()\n)\n\nmean_oof_f1 = float(\n    oof_metrics[\"f1_at_0.50\"].mean()\n)\n\nprint(\n    f\"\\nMean OOF ROC-AUC: \"\n    f\"{mean_oof_auc:.4f}\"\n)\n\nprint(\n    f\"Mean OOF F1 @ 0.50: \"\n    f\"{mean_oof_f1:.4f}\"\n)\n\n# ----------------------------------------------------------------\n# SAVE OOF RESULTS\n# ----------------------------------------------------------------\n\noof_path = os.path.join(\n    AUDIT_DIR,\n    \"cell46_oof_predictions.csv\"\n)\n\nmetrics_path = os.path.join(\n    AUDIT_DIR,\n    \"cell46_oof_metrics.csv\"\n)\n\noof_df.to_csv(\n    oof_path,\n    index=False\n)\n\noof_metrics.to_csv(\n    metrics_path,\n    index=False\n)\n\nprint(\"\\nSaved:\")\nprint(oof_path)\nprint(metrics_path)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 46 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"All 58 labeled studies have exactly one \"\n    \"out-of-fold prediction.\"\n)\n\nprint(\n    \"No training performed in this cell.\"\n)\n\nprint(\n    \"No dataset pipeline modified.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:21:33.969385Z","iopub.execute_input":"2026-08-11T10:21:33.969772Z","iopub.status.idle":"2026-08-11T10:22:04.344042Z","shell.execute_reply.started":"2026-08-11T10:21:33.969741Z","shell.execute_reply":"2026-08-11T10:22:04.342912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 47 - OOF THRESHOLD + CALIBRATION AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    f1_score,\n    precision_score,\n    recall_score,\n    confusion_matrix\n)\n\nprint(\"=\" * 70)\nprint(\"CELL 47 - OOF THRESHOLD + CALIBRATION AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONSTANTS\n# ----------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nOOF_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell46_oof_predictions.csv\"\n)\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\n# ----------------------------------------------------------------\n# LOAD OOF PREDICTIONS\n# ----------------------------------------------------------------\n\nprint(\"\\nLoading OOF predictions...\")\n\nif not os.path.exists(OOF_PATH):\n    raise RuntimeError(\n        \"Missing OOF prediction file: \"\n        + OOF_PATH\n    )\n\noof_df = pd.read_csv(OOF_PATH)\n\nprint(\n    \"OOF table shape:\",\n    oof_df.shape\n)\n\n# ----------------------------------------------------------------\n# BASIC VALIDATION\n# ----------------------------------------------------------------\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"fold\"\n]\n\nfor target in TARGETS:\n    required_columns.append(\n        f\"{target}_label\"\n    )\n    required_columns.append(\n        f\"{target}_prob\"\n    )\n\nmissing_columns = [\n    column\n    for column in required_columns\n    if column not in oof_df.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Missing OOF columns: \"\n        + \", \".join(missing_columns)\n    )\n\nif len(oof_df) != 58:\n    raise RuntimeError(\n        f\"Expected 58 OOF rows, found {len(oof_df)}.\"\n    )\n\nif (\n    oof_df[\"StudyInstanceUID\"].nunique()\n    != 58\n):\n    raise RuntimeError(\n        \"OOF table does not contain exactly \"\n        \"58 unique studies.\"\n    )\n\nif (\n    oof_df.groupby(\"StudyInstanceUID\")\n    .size()\n    .max()\n    != 1\n):\n    raise RuntimeError(\n        \"Duplicate OOF study predictions detected.\"\n    )\n\nprint(\"OOF row count: PASS\")\nprint(\"Unique study count: PASS\")\nprint(\"One prediction per study: PASS\")\n\n# ----------------------------------------------------------------\n# FOLD VALIDATION\n# ----------------------------------------------------------------\n\nfold_counts = (\n    oof_df[\"fold\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(\"\\nOOF fold distribution:\")\nprint(fold_counts)\n\nexpected_fold_counts = {\n    0: 20,\n    1: 19,\n    2: 19\n}\n\nfor fold, expected_count in expected_fold_counts.items():\n\n    actual_count = int(\n        fold_counts.get(fold, 0)\n    )\n\n    if actual_count != expected_count:\n        raise RuntimeError(\n            f\"Fold {fold}: expected \"\n            f\"{expected_count}, found {actual_count}.\"\n        )\n\nprint(\"OOF fold coverage: PASS\")\n\n# ----------------------------------------------------------------\n# PROBABILITY VALIDATION\n# ----------------------------------------------------------------\n\nfor target in TARGETS:\n\n    probabilities = oof_df[\n        f\"{target}_prob\"\n    ].values\n\n    labels = oof_df[\n        f\"{target}_label\"\n    ].values\n\n    if not np.isfinite(probabilities).all():\n        raise RuntimeError(\n            f\"NaN/Inf probabilities detected: {target}\"\n        )\n\n    if not np.isfinite(labels).all():\n        raise RuntimeError(\n            f\"NaN/Inf labels detected: {target}\"\n        )\n\n    if (\n        probabilities.min() < 0.0\n        or probabilities.max() > 1.0\n    ):\n        raise RuntimeError(\n            f\"Invalid probability range: {target}\"\n        )\n\n    unique_labels = set(\n        np.unique(labels).tolist()\n    )\n\n    if not unique_labels.issubset({0.0, 1.0}):\n        raise RuntimeError(\n            f\"Invalid labels for target: {target}\"\n        )\n\nprint(\"Probability validity: PASS\")\nprint(\"Label validity: PASS\")\n\n# ----------------------------------------------------------------\n# THRESHOLD GRID\n# ----------------------------------------------------------------\n\nthresholds = np.round(\n    np.arange(\n        0.10,\n        0.901,\n        0.025\n    ),\n    3\n)\n\nprint(\n    \"\\nThreshold grid:\",\n    f\"{thresholds[0]:.3f}\",\n    \"to\",\n    f\"{thresholds[-1]:.3f}\",\n    f\"({len(thresholds)} thresholds)\"\n)\n\n# ----------------------------------------------------------------\n# TARGET-WISE OOF ANALYSIS\n# ----------------------------------------------------------------\n\nsummary_rows = []\nthreshold_rows = []\n\nfor target in TARGETS:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].astype(int).values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].astype(float).values\n\n    positive_count = int(\n        y_true.sum()\n    )\n\n    negative_count = int(\n        len(y_true) - positive_count\n    )\n\n    # ------------------------------------------------------------\n    # ROC-AUC\n    # ------------------------------------------------------------\n\n    if (\n        len(np.unique(y_true))\n        == 2\n    ):\n        roc_auc = float(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n\n        pr_auc = float(\n            average_precision_score(\n                y_true,\n                y_prob\n            )\n        )\n\n    else:\n        roc_auc = np.nan\n        pr_auc = np.nan\n\n    # ------------------------------------------------------------\n    # BASELINE THRESHOLD\n    # ------------------------------------------------------------\n\n    baseline_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    baseline_f1 = float(\n        f1_score(\n            y_true,\n            baseline_pred,\n            zero_division=0\n        )\n    )\n\n    baseline_precision = float(\n        precision_score(\n            y_true,\n            baseline_pred,\n            zero_division=0\n        )\n    )\n\n    baseline_recall = float(\n        recall_score(\n            y_true,\n            baseline_pred,\n            zero_division=0\n        )\n    )\n\n    # ------------------------------------------------------------\n    # BEST F1 THRESHOLD\n    # ------------------------------------------------------------\n\n    best_threshold = 0.50\n    best_f1 = -1.0\n    best_precision = 0.0\n    best_recall = 0.0\n\n    for threshold in thresholds:\n\n        prediction = (\n            y_prob >= threshold\n        ).astype(int)\n\n        current_f1 = float(\n            f1_score(\n                y_true,\n                prediction,\n                zero_division=0\n            )\n        )\n\n        current_precision = float(\n            precision_score(\n                y_true,\n                prediction,\n                zero_division=0\n            )\n        )\n\n        current_recall = float(\n            recall_score(\n                y_true,\n                prediction,\n                zero_division=0\n            )\n        )\n\n        threshold_rows.append(\n            {\n                \"target\": target,\n                \"threshold\": float(threshold),\n                \"f1\": current_f1,\n                \"precision\": current_precision,\n                \"recall\": current_recall,\n                \"predicted_positive\": int(\n                    prediction.sum()\n                )\n            }\n        )\n\n        if current_f1 > best_f1:\n\n            best_f1 = current_f1\n            best_threshold = float(threshold)\n            best_precision = current_precision\n            best_recall = current_recall\n\n    # ------------------------------------------------------------\n    # PROBABILITY DISTRIBUTION\n    # ------------------------------------------------------------\n\n    positive_probabilities = y_prob[\n        y_true == 1\n    ]\n\n    negative_probabilities = y_prob[\n        y_true == 0\n    ]\n\n    positive_mean = float(\n        positive_probabilities.mean()\n    )\n\n    negative_mean = float(\n        negative_probabilities.mean()\n    )\n\n    probability_separation = (\n        positive_mean\n        - negative_mean\n    )\n\n    # ------------------------------------------------------------\n    # BRIER SCORE\n    # ------------------------------------------------------------\n\n    brier_score = float(\n        np.mean(\n            (\n                y_prob\n                - y_true\n            ) ** 2\n        )\n    )\n\n    # ------------------------------------------------------------\n    # SUMMARY\n    # ------------------------------------------------------------\n\n    summary_rows.append(\n        {\n            \"target\": target,\n            \"positive_count\": positive_count,\n            \"negative_count\": negative_count,\n            \"roc_auc\": roc_auc,\n            \"pr_auc\": pr_auc,\n            \"brier_score\": brier_score,\n            \"f1_at_0.50\": baseline_f1,\n            \"best_threshold\": best_threshold,\n            \"best_f1\": best_f1,\n            \"f1_improvement\": (\n                best_f1\n                - baseline_f1\n            ),\n            \"best_precision\": best_precision,\n            \"best_recall\": best_recall,\n            \"positive_mean_probability\": positive_mean,\n            \"negative_mean_probability\": negative_mean,\n            \"probability_separation\": probability_separation\n        }\n    )\n\n# ----------------------------------------------------------------\n# DATAFRAMES\n# ----------------------------------------------------------------\n\noof_threshold_summary = pd.DataFrame(\n    summary_rows\n)\n\noof_threshold_grid = pd.DataFrame(\n    threshold_rows\n)\n\n# ----------------------------------------------------------------\n# PRINT TARGET SUMMARY\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OOF TARGET THRESHOLD SUMMARY\")\nprint(\"=\" * 70)\n\ndisplay_columns = [\n    \"target\",\n    \"positive_count\",\n    \"negative_count\",\n    \"roc_auc\",\n    \"pr_auc\",\n    \"brier_score\",\n    \"f1_at_0.50\",\n    \"best_threshold\",\n    \"best_f1\",\n    \"f1_improvement\",\n    \"probability_separation\"\n]\n\nprint(\n    oof_threshold_summary[\n        display_columns\n    ].to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# OVERALL SUMMARY\n# ----------------------------------------------------------------\n\nmean_roc_auc = float(\n    oof_threshold_summary[\n        \"roc_auc\"\n    ].mean()\n)\n\nmean_pr_auc = float(\n    oof_threshold_summary[\n        \"pr_auc\"\n    ].mean()\n)\n\nmean_brier = float(\n    oof_threshold_summary[\n        \"brier_score\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    oof_threshold_summary[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nmean_best_f1 = float(\n    oof_threshold_summary[\n        \"best_f1\"\n    ].mean()\n)\n\nmean_f1_improvement = float(\n    oof_threshold_summary[\n        \"f1_improvement\"\n    ].mean()\n)\n\nmean_separation = float(\n    oof_threshold_summary[\n        \"probability_separation\"\n    ].mean()\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OVERALL OOF THRESHOLD SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Mean ROC-AUC:              {mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"Mean PR-AUC:               {mean_pr_auc:.4f}\"\n)\n\nprint(\n    f\"Mean Brier score:          {mean_brier:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50:            {mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1:    {mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement:  {mean_f1_improvement:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: {mean_separation:.4f}\"\n)\n\n# ----------------------------------------------------------------\n# IDENTIFY STRONG / WEAK TARGETS\n# ----------------------------------------------------------------\n\nauc_strong = oof_threshold_summary[\n    oof_threshold_summary[\"roc_auc\"] >= 0.65\n][\n    \"target\"\n].tolist()\n\nauc_weak = oof_threshold_summary[\n    oof_threshold_summary[\"roc_auc\"] < 0.50\n][\n    \"target\"\n].tolist()\n\nnegative_separation = oof_threshold_summary[\n    oof_threshold_summary[\n        \"probability_separation\"\n    ] < 0\n][\n    \"target\"\n].tolist()\n\nlarge_threshold_gain = oof_threshold_summary[\n    oof_threshold_summary[\n        \"f1_improvement\"\n    ] >= 0.10\n][\n    \"target\"\n].tolist()\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OOF TARGET FLAGS\")\nprint(\"=\" * 70)\n\nprint(\n    \"Targets with ROC-AUC >= 0.65:\",\n    len(auc_strong)\n)\n\nprint(\n    auc_strong\n)\n\nprint(\n    \"\\nTargets with ROC-AUC < 0.50:\",\n    len(auc_weak)\n)\n\nprint(\n    auc_weak\n)\n\nprint(\n    \"\\nTargets with negative probability separation:\",\n    len(negative_separation)\n)\n\nprint(\n    negative_separation\n)\n\nprint(\n    \"\\nTargets with >= 0.10 F1 improvement:\",\n    len(large_threshold_gain)\n)\n\nprint(\n    large_threshold_gain\n)\n\n# ----------------------------------------------------------------\n# OOF CONFUSION MATRICES AT 0.50\n# ----------------------------------------------------------------\n\nconfusion_rows = []\n\nfor target in TARGETS:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].astype(int).values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].astype(float).values\n\n    y_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    tn, fp, fn, tp = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=[0, 1]\n    ).ravel()\n\n    sensitivity = (\n        tp / (tp + fn)\n        if (tp + fn) > 0\n        else 0.0\n    )\n\n    specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else 0.0\n    )\n\n    confusion_rows.append(\n        {\n            \"target\": target,\n            \"TP\": int(tp),\n            \"TN\": int(tn),\n            \"FP\": int(fp),\n            \"FN\": int(fn),\n            \"sensitivity\": sensitivity,\n            \"specificity\": specificity\n        }\n    )\n\noof_confusion = pd.DataFrame(\n    confusion_rows\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OOF CONFUSION MATRICES @ 0.50\")\nprint(\"=\" * 70)\n\nprint(\n    oof_confusion.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# SAVE RESULTS\n# ----------------------------------------------------------------\n\nsummary_path = os.path.join(\n    AUDIT_DIR,\n    \"cell47_oof_threshold_summary.csv\"\n)\n\ngrid_path = os.path.join(\n    AUDIT_DIR,\n    \"cell47_oof_threshold_grid.csv\"\n)\n\nconfusion_path = os.path.join(\n    AUDIT_DIR,\n    \"cell47_oof_confusion.csv\"\n)\n\noof_threshold_summary.to_csv(\n    summary_path,\n    index=False\n)\n\noof_threshold_grid.to_csv(\n    grid_path,\n    index=False\n)\n\noof_confusion.to_csv(\n    confusion_path,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL CHECKS\n# ----------------------------------------------------------------\n\nif len(oof_threshold_summary) != 12:\n    raise RuntimeError(\n        \"Expected 12 target summary rows.\"\n    )\n\nif len(oof_threshold_grid) != (\n    12 * len(thresholds)\n):\n    raise RuntimeError(\n        \"Unexpected threshold-grid size.\"\n    )\n\nif len(oof_confusion) != 12:\n    raise RuntimeError(\n        \"Expected 12 confusion-matrix rows.\"\n    )\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 47 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"OOF threshold analysis: PASS\")\nprint(\"OOF calibration/error statistics: PASS\")\nprint(\"No training performed.\")\nprint(\"No model modified.\")\nprint(\"No test data used.\")\n\nprint(\"\\nSaved:\")\nprint(summary_path)\nprint(grid_path)\nprint(confusion_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:27:05.428572Z","iopub.execute_input":"2026-08-11T10:27:05.428947Z","iopub.status.idle":"2026-08-11T10:27:07.631098Z","shell.execute_reply.started":"2026-08-11T10:27:05.428905Z","shell.execute_reply":"2026-08-11T10:27:07.629968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 48 - OOF ERROR + FOLD GENERALIZATION AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    f1_score,\n    precision_score,\n    recall_score\n)\n\nprint(\"=\" * 70)\nprint(\"CELL 48 - OOF ERROR + FOLD GENERALIZATION AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# CONSTANTS\n# ----------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nOOF_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell46_oof_predictions.csv\"\n)\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\"\n]\n\n# ----------------------------------------------------------------\n# LOAD OOF DATA\n# ----------------------------------------------------------------\n\nprint(\"\\nLoading OOF predictions...\")\n\nif not os.path.exists(OOF_PATH):\n    raise RuntimeError(\n        \"Missing OOF prediction file: \"\n        + OOF_PATH\n    )\n\noof_df = pd.read_csv(\n    OOF_PATH\n)\n\nprint(\n    \"OOF shape:\",\n    oof_df.shape\n)\n\n# ----------------------------------------------------------------\n# VALIDATE OOF DATA\n# ----------------------------------------------------------------\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"fold\"\n]\n\nfor target in TARGETS:\n    required_columns.extend([\n        f\"{target}_label\",\n        f\"{target}_prob\"\n    ])\n\nmissing_columns = [\n    column\n    for column in required_columns\n    if column not in oof_df.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Missing OOF columns: \"\n        + \", \".join(missing_columns)\n    )\n\nif len(oof_df) != 58:\n    raise RuntimeError(\n        f\"Expected 58 OOF rows, found {len(oof_df)}.\"\n    )\n\nif (\n    oof_df[\"StudyInstanceUID\"].nunique()\n    != 58\n):\n    raise RuntimeError(\n        \"Expected 58 unique studies.\"\n    )\n\nif (\n    oof_df.groupby(\"StudyInstanceUID\")\n    .size()\n    .max()\n    != 1\n):\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID detected.\"\n    )\n\nif set(\n    oof_df[\"fold\"].astype(int).unique()\n) != {0, 1, 2}:\n\n    raise RuntimeError(\n        \"Expected folds 0, 1, and 2.\"\n    )\n\nprint(\"OOF structure: PASS\")\nprint(\"Study uniqueness: PASS\")\nprint(\"Fold structure: PASS\")\n\n# ----------------------------------------------------------------\n# TARGET-WISE ERROR ANALYSIS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TARGET-WISE OOF ERROR ANALYSIS\")\nprint(\"=\" * 70)\n\ntarget_error_rows = []\n\nfor target in TARGETS:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].astype(int).values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].astype(float).values\n\n    y_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    tp = int(\n        ((y_true == 1) & (y_pred == 1)).sum()\n    )\n\n    tn = int(\n        ((y_true == 0) & (y_pred == 0)).sum()\n    )\n\n    fp = int(\n        ((y_true == 0) & (y_pred == 1)).sum()\n    )\n\n    fn = int(\n        ((y_true == 1) & (y_pred == 0)).sum()\n    )\n\n    error_count = fp + fn\n\n    error_rate = (\n        error_count / len(y_true)\n    )\n\n    f1 = float(\n        f1_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    )\n\n    precision = float(\n        precision_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    )\n\n    recall = float(\n        recall_score(\n            y_true,\n            y_pred,\n            zero_division=0\n        )\n    )\n\n    auc = float(\n        roc_auc_score(\n            y_true,\n            y_prob\n        )\n    )\n\n    target_error_rows.append({\n        \"target\": target,\n        \"TP\": tp,\n        \"TN\": tn,\n        \"FP\": fp,\n        \"FN\": fn,\n        \"total_errors\": error_count,\n        \"error_rate\": error_rate,\n        \"precision\": precision,\n        \"recall\": recall,\n        \"f1\": f1,\n        \"roc_auc\": auc\n    })\n\ntarget_error_df = pd.DataFrame(\n    target_error_rows\n)\n\nprint(\n    target_error_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# FOLD-WISE PERFORMANCE\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD-WISE OOF PERFORMANCE\")\nprint(\"=\" * 70)\n\nfold_rows = []\n\nfor fold in [0, 1, 2]:\n\n    fold_df = oof_df[\n        oof_df[\"fold\"].astype(int)\n        == fold\n    ].copy()\n\n    fold_auc_values = []\n    fold_f1_values = []\n\n    for target in TARGETS:\n\n        y_true = fold_df[\n            f\"{target}_label\"\n        ].astype(int).values\n\n        y_prob = fold_df[\n            f\"{target}_prob\"\n        ].astype(float).values\n\n        y_pred = (\n            y_prob >= 0.50\n        ).astype(int)\n\n        # Every fold was explicitly checked for\n        # both classes, but retain a safety check.\n        if len(\n            np.unique(y_true)\n        ) < 2:\n            continue\n\n        fold_auc_values.append(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n\n        fold_f1_values.append(\n            f1_score(\n                y_true,\n                y_pred,\n                zero_division=0\n            )\n        )\n\n    fold_rows.append({\n        \"fold\": fold,\n        \"studies\": len(fold_df),\n        \"mean_roc_auc\": float(\n            np.mean(fold_auc_values)\n        ),\n        \"mean_f1\": float(\n            np.mean(fold_f1_values)\n        )\n    })\n\nfold_performance_df = pd.DataFrame(\n    fold_rows\n)\n\nprint(\n    fold_performance_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# TARGET x FOLD ROC-AUC MATRIX\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TARGET x FOLD ROC-AUC\")\nprint(\"=\" * 70)\n\nauc_matrix = []\n\nfor target in TARGETS:\n\n    row = {\n        \"target\": target\n    }\n\n    for fold in [0, 1, 2]:\n\n        fold_df = oof_df[\n            oof_df[\"fold\"].astype(int)\n            == fold\n        ]\n\n        y_true = fold_df[\n            f\"{target}_label\"\n        ].astype(int).values\n\n        y_prob = fold_df[\n            f\"{target}_prob\"\n        ].astype(float).values\n\n        if len(\n            np.unique(y_true)\n        ) < 2:\n\n            row[\n                f\"fold_{fold}_auc\"\n            ] = np.nan\n\n        else:\n\n            row[\n                f\"fold_{fold}_auc\"\n            ] = float(\n                roc_auc_score(\n                    y_true,\n                    y_prob\n                )\n            )\n\n    auc_matrix.append(row)\n\nauc_matrix_df = pd.DataFrame(\n    auc_matrix\n)\n\nprint(\n    auc_matrix_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# TARGET x FOLD F1 MATRIX\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TARGET x FOLD F1 @ 0.50\")\nprint(\"=\" * 70)\n\nf1_matrix = []\n\nfor target in TARGETS:\n\n    row = {\n        \"target\": target\n    }\n\n    for fold in [0, 1, 2]:\n\n        fold_df = oof_df[\n            oof_df[\"fold\"].astype(int)\n            == fold\n        ]\n\n        y_true = fold_df[\n            f\"{target}_label\"\n        ].astype(int).values\n\n        y_prob = fold_df[\n            f\"{target}_prob\"\n        ].astype(float).values\n\n        y_pred = (\n            y_prob >= 0.50\n        ).astype(int)\n\n        row[\n            f\"fold_{fold}_f1\"\n        ] = float(\n            f1_score(\n                y_true,\n                y_pred,\n                zero_division=0\n            )\n        )\n\n    f1_matrix.append(row)\n\nf1_matrix_df = pd.DataFrame(\n    f1_matrix\n)\n\nprint(\n    f1_matrix_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# CROSS-FOLD STABILITY\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CROSS-FOLD TARGET STABILITY\")\nprint(\"=\" * 70)\n\nstability_rows = []\n\nfor target in TARGETS:\n\n    auc_values = (\n        auc_matrix_df[\n            auc_matrix_df[\"target\"]\n            == target\n        ][\n            [\n                \"fold_0_auc\",\n                \"fold_1_auc\",\n                \"fold_2_auc\"\n            ]\n        ]\n        .values\n        .flatten()\n    )\n\n    f1_values = (\n        f1_matrix_df[\n            f1_matrix_df[\"target\"]\n            == target\n        ][\n            [\n                \"fold_0_f1\",\n                \"fold_1_f1\",\n                \"fold_2_f1\"\n            ]\n        ]\n        .values\n        .flatten()\n    )\n\n    stability_rows.append({\n        \"target\": target,\n        \"auc_mean\": float(\n            np.nanmean(auc_values)\n        ),\n        \"auc_std\": float(\n            np.nanstd(auc_values)\n        ),\n        \"auc_min\": float(\n            np.nanmin(auc_values)\n        ),\n        \"auc_max\": float(\n            np.nanmax(auc_values)\n        ),\n        \"f1_mean\": float(\n            np.mean(f1_values)\n        ),\n        \"f1_std\": float(\n            np.std(f1_values)\n        )\n    })\n\nstability_df = pd.DataFrame(\n    stability_rows\n)\n\nprint(\n    stability_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# STUDY-LEVEL ERROR BURDEN\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STUDY-LEVEL ERROR BURDEN\")\nprint(\"=\" * 70)\n\nstudy_error_rows = []\n\nfor _, row in oof_df.iterrows():\n\n    study_id = row[\n        \"StudyInstanceUID\"\n    ]\n\n    fold = int(\n        row[\"fold\"]\n    )\n\n    total_errors = 0\n    total_positive_targets = 0\n    total_predicted_positive = 0\n    total_high_confidence_errors = 0\n\n    for target in TARGETS:\n\n        label = int(\n            row[f\"{target}_label\"]\n        )\n\n        probability = float(\n            row[f\"{target}_prob\"]\n        )\n\n        prediction = int(\n            probability >= 0.50\n        )\n\n        if label == 1:\n            total_positive_targets += 1\n\n        if prediction == 1:\n            total_predicted_positive += 1\n\n        if prediction != label:\n            total_errors += 1\n\n            # High-confidence error.\n            if (\n                probability >= 0.75\n                or probability <= 0.25\n            ):\n                total_high_confidence_errors += 1\n\n    study_error_rows.append({\n        \"StudyInstanceUID\": study_id,\n        \"fold\": fold,\n        \"total_target_errors\": total_errors,\n        \"positive_targets\": total_positive_targets,\n        \"predicted_positive_targets\": total_predicted_positive,\n        \"high_confidence_errors\": (\n            total_high_confidence_errors\n        )\n    })\n\nstudy_error_df = pd.DataFrame(\n    study_error_rows\n)\n\nstudy_error_df = study_error_df.sort_values(\n    [\n        \"total_target_errors\",\n        \"high_confidence_errors\"\n    ],\n    ascending=False\n).reset_index(drop=True)\n\nprint(\n    \"\\nMost difficult studies by total target errors:\"\n)\n\nprint(\n    study_error_df.head(\n        15\n    ).to_string(\n        index=False\n    )\n)\n\n# ----------------------------------------------------------------\n# HIGH-CONFIDENCE ERROR SUMMARY\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"HIGH-CONFIDENCE ERROR ANALYSIS\")\nprint(\"=\" * 70)\n\nhigh_confidence_error_rows = []\n\nfor target in TARGETS:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].astype(int).values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].astype(float).values\n\n    y_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    high_confidence_mask = (\n        (\n            y_pred != y_true\n        )\n        &\n        (\n            (y_prob >= 0.75)\n            |\n            (y_prob <= 0.25)\n        )\n    )\n\n    high_confidence_error_rows.append({\n        \"target\": target,\n        \"high_confidence_errors\": int(\n            high_confidence_mask.sum()\n        ),\n        \"total_errors\": int(\n            (y_pred != y_true).sum()\n        ),\n        \"high_confidence_error_rate\": float(\n            high_confidence_mask.mean()\n        )\n    })\n\nhigh_confidence_df = pd.DataFrame(\n    high_confidence_error_rows\n)\n\nprint(\n    high_confidence_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# ERROR DIRECTION\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FALSE POSITIVE / FALSE NEGATIVE BURDEN\")\nprint(\"=\" * 70)\n\nerror_direction_rows = []\n\nfor target in TARGETS:\n\n    y_true = oof_df[\n        f\"{target}_label\"\n    ].astype(int).values\n\n    y_prob = oof_df[\n        f\"{target}_prob\"\n    ].astype(float).values\n\n    y_pred = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    fp = int(\n        (\n            (y_true == 0)\n            &\n            (y_pred == 1)\n        ).sum()\n    )\n\n    fn = int(\n        (\n            (y_true == 1)\n            &\n            (y_pred == 0)\n        ).sum()\n    )\n\n    if fp > fn:\n        dominant_error = \"False Positive\"\n    elif fn > fp:\n        dominant_error = \"False Negative\"\n    else:\n        dominant_error = \"Balanced\"\n\n    error_direction_rows.append({\n        \"target\": target,\n        \"false_positives\": fp,\n        \"false_negatives\": fn,\n        \"FP_minus_FN\": fp - fn,\n        \"dominant_error\": dominant_error\n    })\n\nerror_direction_df = pd.DataFrame(\n    error_direction_rows\n)\n\nprint(\n    error_direction_df.to_string(\n        index=False\n    )\n)\n\n# ----------------------------------------------------------------\n# OVERALL MODEL DIAGNOSTIC\n# ----------------------------------------------------------------\n\nmean_auc = float(\n    target_error_df[\n        \"roc_auc\"\n    ].mean()\n)\n\nmean_f1 = float(\n    target_error_df[\n        \"f1\"\n    ].mean()\n)\n\nmean_error_rate = float(\n    target_error_df[\n        \"error_rate\"\n    ].mean()\n)\n\ntotal_fp = int(\n    error_direction_df[\n        \"false_positives\"\n    ].sum()\n)\n\ntotal_fn = int(\n    error_direction_df[\n        \"false_negatives\"\n    ].sum()\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OVERALL ERROR DIAGNOSTIC\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Mean OOF ROC-AUC:       {mean_auc:.4f}\"\n)\n\nprint(\n    f\"Mean OOF F1 @ 0.50:     {mean_f1:.4f}\"\n)\n\nprint(\n    f\"Mean target error rate: {mean_error_rate:.4f}\"\n)\n\nprint(\n    f\"Total false positives:  {total_fp}\"\n)\n\nprint(\n    f\"Total false negatives:  {total_fn}\"\n)\n\nprint(\n    f\"FP/FN ratio:            \"\n    f\"{total_fp / max(total_fn, 1):.4f}\"\n)\n\n# ----------------------------------------------------------------\n# SAVE RESULTS\n# ----------------------------------------------------------------\n\ntarget_error_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_oof_target_errors.csv\"\n)\n\nfold_performance_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_oof_fold_performance.csv\"\n)\n\nauc_matrix_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_target_fold_auc.csv\"\n)\n\nf1_matrix_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_target_fold_f1.csv\"\n)\n\nstability_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_target_stability.csv\"\n)\n\nstudy_error_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_study_error_burden.csv\"\n)\n\nhigh_confidence_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_high_confidence_errors.csv\"\n)\n\nerror_direction_path = os.path.join(\n    AUDIT_DIR,\n    \"cell48_error_direction.csv\"\n)\n\ntarget_error_df.to_csv(\n    target_error_path,\n    index=False\n)\n\nfold_performance_df.to_csv(\n    fold_performance_path,\n    index=False\n)\n\nauc_matrix_df.to_csv(\n    auc_matrix_path,\n    index=False\n)\n\nf1_matrix_df.to_csv(\n    f1_matrix_path,\n    index=False\n)\n\nstability_df.to_csv(\n    stability_path,\n    index=False\n)\n\nstudy_error_df.to_csv(\n    study_error_path,\n    index=False\n)\n\nhigh_confidence_df.to_csv(\n    high_confidence_path,\n    index=False\n)\n\nerror_direction_df.to_csv(\n    error_direction_path,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL VALIDATION\n# ----------------------------------------------------------------\n\nif len(target_error_df) != 12:\n    raise RuntimeError(\n        \"Target error table must contain 12 targets.\"\n    )\n\nif len(fold_performance_df) != 3:\n    raise RuntimeError(\n        \"Fold performance table must contain 3 folds.\"\n    )\n\nif len(auc_matrix_df) != 12:\n    raise RuntimeError(\n        \"AUC matrix must contain 12 targets.\"\n    )\n\nif len(f1_matrix_df) != 12:\n    raise RuntimeError(\n        \"F1 matrix must contain 12 targets.\"\n    )\n\nif len(stability_df) != 12:\n    raise RuntimeError(\n        \"Stability table must contain 12 targets.\"\n    )\n\nif len(study_error_df) != 58:\n    raise RuntimeError(\n        \"Study error table must contain 58 studies.\"\n    )\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 48 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\"Target-wise error analysis: PASS\")\nprint(\"Fold-wise generalization analysis: PASS\")\nprint(\"Study-level error analysis: PASS\")\nprint(\"High-confidence error analysis: PASS\")\nprint(\"No training performed.\")\nprint(\"No model modified.\")\nprint(\"No test data used.\")\n\nprint(\"\\nSaved:\")\nprint(target_error_path)\nprint(fold_performance_path)\nprint(auc_matrix_path)\nprint(f1_matrix_path)\nprint(stability_path)\nprint(study_error_path)\nprint(high_confidence_path)\nprint(error_direction_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:29:14.684195Z","iopub.execute_input":"2026-08-11T10:29:14.684533Z","iopub.status.idle":"2026-08-11T10:29:15.210617Z","shell.execute_reply.started":"2026-08-11T10:29:14.684505Z","shell.execute_reply":"2026-08-11T10:29:15.209625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 49 - MODEL CAPACITY + TRAINABLE LAYER AUDIT\n# ================================================================\n\nimport os\nimport torch\nimport torch.nn as nn\n\nprint(\"=\" * 70)\nprint(\"CELL 49 - MODEL CAPACITY + TRAINABLE LAYER AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECT CHECK\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\"\n]\n\nmissing_objects = []\n\nfor name in required_objects:\n    if name not in globals():\n        missing_objects.append(name)\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE CHECK\n# ----------------------------------------------------------------\n\nmodel_device = next(\n    model.parameters()\n).device\n\nprint(\n    \"Model device:\",\n    model_device\n)\n\n# ----------------------------------------------------------------\n# BASIC MODEL INFORMATION\n# ----------------------------------------------------------------\n\ntotal_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n)\n\ntrainable_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\nfrozen_parameters = (\n    total_parameters\n    - trainable_parameters\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MODEL PARAMETER SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Total parameters:      {total_parameters:,}\"\n)\n\nprint(\n    f\"Trainable parameters:  {trainable_parameters:,}\"\n)\n\nprint(\n    f\"Frozen parameters:     {frozen_parameters:,}\"\n)\n\nprint(\n    f\"Trainable percentage:  \"\n    f\"{100.0 * trainable_parameters / total_parameters:.2f}%\"\n)\n\n# ----------------------------------------------------------------\n# MODEL STRUCTURE\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MODEL STRUCTURE\")\nprint(\"=\" * 70)\n\nprint(model)\n\n# ----------------------------------------------------------------\n# TRAINABLE PARAMETER BREAKDOWN\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINABLE PARAMETER BREAKDOWN\")\nprint(\"=\" * 70)\n\ntrainable_rows = []\n\nfor name, parameter in model.named_parameters():\n\n    if parameter.requires_grad:\n\n        trainable_rows.append({\n            \"parameter\": name,\n            \"shape\": tuple(\n                parameter.shape\n            ),\n            \"num_parameters\": parameter.numel()\n        })\n\nif len(trainable_rows) == 0:\n    raise RuntimeError(\n        \"No trainable parameters found.\"\n    )\n\nfor row in trainable_rows:\n\n    print(\n        f\"{row['parameter']:<50}\"\n        f\"shape={str(row['shape']):<25}\"\n        f\"params={row['num_parameters']:,}\"\n    )\n\ntrainable_df = None\n\n# ----------------------------------------------------------------\n# TRAINABLE PARAMETERS BY TOP-LEVEL MODULE\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINABLE PARAMETERS BY TOP-LEVEL MODULE\")\nprint(\"=\" * 70)\n\nmodule_parameter_counts = {}\n\nfor name, parameter in model.named_parameters():\n\n    if not parameter.requires_grad:\n        continue\n\n    top_level = name.split(\".\")[0]\n\n    if top_level not in module_parameter_counts:\n        module_parameter_counts[\n            top_level\n        ] = 0\n\n    module_parameter_counts[\n        top_level\n    ] += parameter.numel()\n\nfor module_name, count in sorted(\n    module_parameter_counts.items()\n):\n\n    print(\n        f\"{module_name:<25}\"\n        f\"{count:,}\"\n    )\n\n# ----------------------------------------------------------------\n# ALL PARAMETER MODULE STATUS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"MODULE TRAINABILITY STATUS\")\nprint(\"=\" * 70)\n\nmodule_rows = []\n\nfor name, module in model.named_modules():\n\n    if name == \"\":\n        continue\n\n    parameters = list(\n        module.parameters(\n            recurse=False\n        )\n    )\n\n    if not parameters:\n        continue\n\n    module_total = sum(\n        p.numel()\n        for p in parameters\n    )\n\n    module_trainable = sum(\n        p.numel()\n        for p in parameters\n        if p.requires_grad\n    )\n\n    if module_trainable > 0:\n        status = \"TRAINABLE\"\n    else:\n        status = \"FROZEN\"\n\n    module_rows.append({\n        \"module\": name,\n        \"status\": status,\n        \"total_parameters\": module_total,\n        \"trainable_parameters\": module_trainable\n    })\n\n    print(\n        f\"{name:<35}\"\n        f\"{status:<12}\"\n        f\"total={module_total:,} \"\n        f\"trainable={module_trainable:,}\"\n    )\n\n# ----------------------------------------------------------------\n# INPUT CONVOLUTION AUDIT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"INPUT CHANNEL AUDIT\")\nprint(\"=\" * 70)\n\nconv_candidates = []\n\nfor name, module in model.named_modules():\n\n    if isinstance(\n        module,\n        nn.Conv2d\n    ):\n\n        conv_candidates.append(\n            (\n                name,\n                module\n            )\n        )\n\nif len(conv_candidates) == 0:\n    raise RuntimeError(\n        \"No Conv2d layers found.\"\n    )\n\nfirst_conv_name, first_conv = (\n    conv_candidates[0]\n)\n\nprint(\n    \"First Conv2d:\",\n    first_conv_name\n)\n\nprint(\n    \"Input channels:\",\n    first_conv.in_channels\n)\n\nprint(\n    \"Output channels:\",\n    first_conv.out_channels\n)\n\nprint(\n    \"Kernel:\",\n    first_conv.kernel_size\n)\n\nif first_conv.in_channels != 21:\n    raise RuntimeError(\n        \"Expected first convolution to accept \"\n        \"21 input channels, found \"\n        f\"{first_conv.in_channels}.\"\n    )\n\nprint(\n    \"21-channel input: PASS\"\n)\n\n# ----------------------------------------------------------------\n# OUTPUT HEAD AUDIT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OUTPUT HEAD AUDIT\")\nprint(\"=\" * 70)\n\nlinear_candidates = []\n\nfor name, module in model.named_modules():\n\n    if isinstance(\n        module,\n        nn.Linear\n    ):\n\n        linear_candidates.append(\n            (\n                name,\n                module\n            )\n        )\n\nif len(linear_candidates) == 0:\n    raise RuntimeError(\n        \"No Linear layers found.\"\n    )\n\nfor name, module in linear_candidates:\n\n    print(\n        f\"{name}: \"\n        f\"in_features={module.in_features}, \"\n        f\"out_features={module.out_features}, \"\n        f\"trainable={any(p.requires_grad for p in module.parameters())}\"\n    )\n\nfinal_linear_name, final_linear = (\n    linear_candidates[-1]\n)\n\nif final_linear.out_features != len(TARGETS):\n    raise RuntimeError(\n        \"Output head mismatch: expected \"\n        f\"{len(TARGETS)} outputs, found \"\n        f\"{final_linear.out_features}.\"\n    )\n\nprint(\n    f\"Output classes: {final_linear.out_features}\"\n)\n\nprint(\n    \"12-target output head: PASS\"\n)\n\n# ----------------------------------------------------------------\n# TRAINABLE PARAMETER COUNT CONSISTENCY\n# ----------------------------------------------------------------\n\nrecomputed_trainable = sum(\n    row[\"num_parameters\"]\n    for row in trainable_rows\n)\n\nif (\n    recomputed_trainable\n    != trainable_parameters\n):\n    raise RuntimeError(\n        \"Trainable parameter count mismatch.\"\n    )\n\nprint(\n    \"\\nTrainable parameter count consistency: PASS\"\n)\n\n# ----------------------------------------------------------------\n# MODEL FINITE PARAMETER CHECK\n# ----------------------------------------------------------------\n\ninvalid_parameters = []\n\nfor name, parameter in model.named_parameters():\n\n    if not torch.isfinite(\n        parameter.detach()\n    ).all():\n\n        invalid_parameters.append(\n            name\n        )\n\nif invalid_parameters:\n\n    raise RuntimeError(\n        \"NaN/Inf detected in model parameters: \"\n        + \", \".join(\n            invalid_parameters\n        )\n    )\n\nprint(\n    \"Model parameter NaN/Inf check: PASS\"\n)\n\n# ----------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 49 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    \"Model structure inspection: PASS\"\n)\n\nprint(\n    \"Trainable/frozen parameter audit: PASS\"\n)\n\nprint(\n    \"21-channel input configuration: PASS\"\n)\n\nprint(\n    \"12-target multilabel output: PASS\"\n)\n\nprint(\n    \"Model parameter validity: PASS\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No model parameters modified.\"\n)\n\nprint(\n    \"No dataset modified.\"\n)\n\nprint(\n    \"No validation/test data modified.\"\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 49 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"\\nUse this output to determine the next \"\n    \"controlled fine-tuning configuration.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:31:14.891427Z","iopub.execute_input":"2026-08-11T10:31:14.891782Z","iopub.status.idle":"2026-08-11T10:31:14.967068Z","shell.execute_reply.started":"2026-08-11T10:31:14.891753Z","shell.execute_reply":"2026-08-11T10:31:14.965955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 50 - CONTROLLED FINE-TUNING CONFIGURATION AUDIT\n# ================================================================\n\nimport torch\nimport torch.nn as nn\n\nprint(\"=\" * 70)\nprint(\"CELL 50 - CONTROLLED FINE-TUNING CONFIGURATION AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECT CHECK\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = next(\n    model.parameters()\n).device\n\nprint(\n    \"Device:\",\n    device\n)\n\n# ----------------------------------------------------------------\n# VERIFY MODEL STRUCTURE\n# ----------------------------------------------------------------\n\nif not hasattr(model, \"backbone\"):\n    raise RuntimeError(\n        \"Model does not contain 'backbone'.\"\n    )\n\nif not hasattr(model, \"classifier\"):\n    raise RuntimeError(\n        \"Model does not contain 'classifier'.\"\n    )\n\nif not hasattr(model.backbone, \"layer4\"):\n    raise RuntimeError(\n        \"Backbone does not contain 'layer4'.\"\n    )\n\nprint(\n    \"Backbone structure: PASS\"\n)\n\nprint(\n    \"Layer4 availability: PASS\"\n)\n\nprint(\n    \"Classifier availability: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY INPUT CHANNELS\n# ----------------------------------------------------------------\n\nfirst_conv = model.backbone.conv1\n\nif first_conv.in_channels != 21:\n    raise RuntimeError(\n        \"Expected 21 input channels, found \"\n        f\"{first_conv.in_channels}.\"\n    )\n\nprint(\n    \"21-channel input: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY OUTPUT CLASSES\n# ----------------------------------------------------------------\n\nlinear_layers = [\n    module\n    for module in model.classifier.modules()\n    if isinstance(module, nn.Linear)\n]\n\nif not linear_layers:\n    raise RuntimeError(\n        \"No Linear layers found in classifier.\"\n    )\n\nfinal_classifier = linear_layers[-1]\n\nif final_classifier.out_features != len(TARGETS):\n    raise RuntimeError(\n        \"Classifier output mismatch. Expected \"\n        f\"{len(TARGETS)}, found \"\n        f\"{final_classifier.out_features}.\"\n    )\n\nprint(\n    f\"12-target output: PASS\"\n)\n\n# ----------------------------------------------------------------\n# RESET TRAINABILITY\n#\n# Important:\n# We only change requires_grad flags here.\n# No parameter values are modified.\n# ----------------------------------------------------------------\n\nfor parameter in model.parameters():\n    parameter.requires_grad = False\n\n# Keep the 21-channel input adaptation trainable.\nfor parameter in model.backbone.conv1.parameters():\n    parameter.requires_grad = True\n\n# Unfreeze only the deepest ResNet block.\nfor parameter in model.backbone.layer4.parameters():\n    parameter.requires_grad = True\n\n# Keep the classification head trainable.\nfor parameter in model.classifier.parameters():\n    parameter.requires_grad = True\n\nprint(\n    \"\\nTrainability configuration applied.\"\n)\n\n# ----------------------------------------------------------------\n# TRAINABLE MODULE VERIFICATION\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINABILITY VERIFICATION\")\nprint(\"=\" * 70)\n\nexpected_trainable_prefixes = [\n    \"backbone.conv1\",\n    \"backbone.layer4\",\n    \"classifier\"\n]\n\nunexpected_trainable = []\n\nfor name, parameter in model.named_parameters():\n\n    should_train = any(\n        name.startswith(prefix)\n        for prefix in expected_trainable_prefixes\n    )\n\n    if parameter.requires_grad != should_train:\n        unexpected_trainable.append(\n            name\n        )\n\nif unexpected_trainable:\n    raise RuntimeError(\n        \"Unexpected trainability state detected:\\n\"\n        + \"\\n\".join(unexpected_trainable)\n    )\n\nprint(\n    \"Expected trainable modules: PASS\"\n)\n\n# ----------------------------------------------------------------\n# PARAMETER COUNTS\n# ----------------------------------------------------------------\n\ntotal_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n)\n\ntrainable_parameters = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\nfrozen_parameters = (\n    total_parameters\n    - trainable_parameters\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE-TUNING PARAMETER SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Total parameters:      {total_parameters:,}\"\n)\n\nprint(\n    f\"Trainable parameters:  {trainable_parameters:,}\"\n)\n\nprint(\n    f\"Frozen parameters:     {frozen_parameters:,}\"\n)\n\nprint(\n    f\"Trainable percentage:  \"\n    f\"{100.0 * trainable_parameters / total_parameters:.2f}%\"\n)\n\n# ----------------------------------------------------------------\n# MODULE-WISE TRAINABLE COUNTS\n# ----------------------------------------------------------------\n\nmodule_counts = {}\n\nfor name, parameter in model.named_parameters():\n\n    if not parameter.requires_grad:\n        continue\n\n    top_level = name.split(\".\")[0]\n\n    module_counts[top_level] = (\n        module_counts.get(top_level, 0)\n        + parameter.numel()\n    )\n\nprint(\"\\nTrainable parameters by top-level module:\")\n\nfor name, count in sorted(\n    module_counts.items()\n):\n    print(\n        f\"{name:<20} {count:,}\"\n    )\n\n# ----------------------------------------------------------------\n# LAYER4 COUNT\n# ----------------------------------------------------------------\n\nlayer4_trainable = sum(\n    parameter.numel()\n    for parameter in model.backbone.layer4.parameters()\n    if parameter.requires_grad\n)\n\nconv1_trainable = sum(\n    parameter.numel()\n    for parameter in model.backbone.conv1.parameters()\n    if parameter.requires_grad\n)\n\nclassifier_trainable = sum(\n    parameter.numel()\n    for parameter in model.classifier.parameters()\n    if parameter.requires_grad\n)\n\nprint(\"\\nDetailed trainable counts:\")\n\nprint(\n    f\"Conv1:       {conv1_trainable:,}\"\n)\n\nprint(\n    f\"Layer4:      {layer4_trainable:,}\"\n)\n\nprint(\n    f\"Classifier:  {classifier_trainable:,}\"\n)\n\nexpected_total = (\n    conv1_trainable\n    + layer4_trainable\n    + classifier_trainable\n)\n\nif trainable_parameters != expected_total:\n    raise RuntimeError(\n        \"Trainable parameter count does not match \"\n        \"Conv1 + Layer4 + Classifier.\"\n    )\n\nprint(\n    \"Trainable parameter arithmetic: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY EARLIER RESNET LAYERS REMAIN FROZEN\n# ----------------------------------------------------------------\n\nearlier_layers = [\n    \"bn1\",\n    \"layer1\",\n    \"layer2\",\n    \"layer3\"\n]\n\nunexpected_earlier_training = []\n\nfor layer_name in earlier_layers:\n\n    layer = getattr(\n        model.backbone,\n        layer_name\n    )\n\n    if any(\n        parameter.requires_grad\n        for parameter in layer.parameters()\n    ):\n        unexpected_earlier_training.append(\n            layer_name\n        )\n\nif unexpected_earlier_training:\n    raise RuntimeError(\n        \"Unexpectedly trainable earlier layers: \"\n        + \", \".join(\n            unexpected_earlier_training\n        )\n    )\n\nprint(\n    \"Backbone layers 1-3 frozen: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY CLASSIFIER\n# ----------------------------------------------------------------\n\nif not all(\n    parameter.requires_grad\n    for parameter in model.classifier.parameters()\n):\n    raise RuntimeError(\n        \"Classifier is not fully trainable.\"\n    )\n\nprint(\n    \"Classifier trainability: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY MODEL PARAMETERS ARE FINITE\n# ----------------------------------------------------------------\n\ninvalid_parameters = []\n\nfor name, parameter in model.named_parameters():\n\n    if not torch.isfinite(\n        parameter.detach()\n    ).all():\n\n        invalid_parameters.append(\n            name\n        )\n\nif invalid_parameters:\n    raise RuntimeError(\n        \"NaN/Inf detected in model parameters: \"\n        + \", \".join(invalid_parameters)\n    )\n\nprint(\n    \"Model parameter validity: PASS\"\n)\n\n# ----------------------------------------------------------------\n# OPTIMIZER PLAN\n#\n# Different learning rates are intentional.\n# Layer4 needs a small update.\n# Classifier and the new 21-channel conv1 adaptation\n# can learn faster.\n# ----------------------------------------------------------------\n\nFINE_TUNE_LR_BACKBONE = 1e-5\nFINE_TUNE_LR_CONV1 = 1e-4\nFINE_TUNE_LR_CLASSIFIER = 1e-3\nFINE_TUNE_WEIGHT_DECAY = 1e-4\nFINE_TUNE_GRAD_CLIP = 1.0\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CONTROLLED FINE-TUNING CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Conv1 learning rate:       {FINE_TUNE_LR_CONV1}\"\n)\n\nprint(\n    f\"Layer4 learning rate:      {FINE_TUNE_LR_BACKBONE}\"\n)\n\nprint(\n    f\"Classifier learning rate:  {FINE_TUNE_LR_CLASSIFIER}\"\n)\n\nprint(\n    f\"Weight decay:               {FINE_TUNE_WEIGHT_DECAY}\"\n)\n\nprint(\n    f\"Gradient clipping:          {FINE_TUNE_GRAD_CLIP}\"\n)\n\n# ----------------------------------------------------------------\n# CREATE OPTIMIZER\n#\n# Only trainable parameters are included.\n# ----------------------------------------------------------------\n\noptimizer_param_groups = [\n    {\n        \"params\": model.backbone.conv1.parameters(),\n        \"lr\": FINE_TUNE_LR_CONV1\n    },\n    {\n        \"params\": model.backbone.layer4.parameters(),\n        \"lr\": FINE_TUNE_LR_BACKBONE\n    },\n    {\n        \"params\": model.classifier.parameters(),\n        \"lr\": FINE_TUNE_LR_CLASSIFIER\n    }\n]\n\nfine_tune_optimizer = torch.optim.AdamW(\n    optimizer_param_groups,\n    weight_decay=FINE_TUNE_WEIGHT_DECAY\n)\n\nprint(\n    \"\\nFine-tuning optimizer created: PASS\"\n)\n\nprint(\n    \"Optimizer: AdamW\"\n)\n\nprint(\n    \"Parameter groups: 3\"\n)\n\n# ----------------------------------------------------------------\n# OPTIMIZER PARAMETER COUNT\n# ----------------------------------------------------------------\n\noptimizer_parameter_count = sum(\n    parameter.numel()\n    for group in fine_tune_optimizer.param_groups\n    for parameter in group[\"params\"]\n)\n\nif optimizer_parameter_count != trainable_parameters:\n    raise RuntimeError(\n        \"Optimizer parameter count mismatch. \"\n        f\"Optimizer={optimizer_parameter_count:,}, \"\n        f\"Trainable={trainable_parameters:,}\"\n    )\n\nprint(\n    \"Optimizer/trainable parameter count: PASS\"\n)\n\n# ----------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 50 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    \"Controlled Layer4 fine-tuning configuration: READY\"\n)\n\nprint(\n    \"Conv1 trainable: YES\"\n)\n\nprint(\n    \"Layer1 frozen: YES\"\n)\n\nprint(\n    \"Layer2 frozen: YES\"\n)\n\nprint(\n    \"Layer3 frozen: YES\"\n)\n\nprint(\n    \"Layer4 trainable: YES\"\n)\n\nprint(\n    \"Classifier trainable: YES\"\n)\n\nprint(\n    \"Differential learning rates: READY\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No checkpoint overwritten.\"\n)\n\nprint(\n    \"No dataset modified.\"\n)\n\nprint(\n    \"No validation/test predictions generated.\"\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 50 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"\\nNext step: controlled fine-tuning smoke test.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:35:34.380301Z","iopub.execute_input":"2026-08-11T10:35:34.380648Z","iopub.status.idle":"2026-08-11T10:35:34.458955Z","shell.execute_reply.started":"2026-08-11T10:35:34.380618Z","shell.execute_reply":"2026-08-11T10:35:34.458085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 51 - CONTROLLED FINE-TUNING SMOKE TEST\n# ================================================================\n\nimport copy\nimport torch\n\nprint(\"=\" * 70)\nprint(\"CELL 51 - CONTROLLED FINE-TUNING SMOKE TEST\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"train_loader\",\n    \"TARGETS\",\n    \"fine_tune_optimizer\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = next(model.parameters()).device\n\nprint(\n    f\"Device: {device}\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY TRAINABLE PARAMETER CONFIGURATION\n# ----------------------------------------------------------------\n\nexpected_trainable_prefixes = [\n    \"backbone.conv1\",\n    \"backbone.layer4\",\n    \"classifier\"\n]\n\nunexpected_trainable = []\n\nfor name, parameter in model.named_parameters():\n\n    expected = any(\n        name.startswith(prefix)\n        for prefix in expected_trainable_prefixes\n    )\n\n    if parameter.requires_grad != expected:\n        unexpected_trainable.append(name)\n\nif unexpected_trainable:\n    raise RuntimeError(\n        \"Unexpected trainability state:\\n\"\n        + \"\\n\".join(unexpected_trainable)\n    )\n\nprint(\n    \"Trainability configuration: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY FROZEN LAYERS\n# ----------------------------------------------------------------\n\nfrozen_prefixes = [\n    \"backbone.bn1\",\n    \"backbone.layer1\",\n    \"backbone.layer2\",\n    \"backbone.layer3\"\n]\n\nfor prefix in frozen_prefixes:\n\n    matching = [\n        parameter\n        for name, parameter in model.named_parameters()\n        if name.startswith(prefix)\n    ]\n\n    if not matching:\n        raise RuntimeError(\n            f\"No parameters found for frozen module: {prefix}\"\n        )\n\n    if any(\n        parameter.requires_grad\n        for parameter in matching\n    ):\n        raise RuntimeError(\n            f\"Frozen module became trainable: {prefix}\"\n        )\n\nprint(\n    \"Frozen backbone layers 1-3: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY OPTIMIZER\n# ----------------------------------------------------------------\n\noptimizer_parameter_ids = {\n    id(parameter)\n    for group in fine_tune_optimizer.param_groups\n    for parameter in group[\"params\"]\n}\n\ntrainable_parameter_ids = {\n    id(parameter)\n    for parameter in model.parameters()\n    if parameter.requires_grad\n}\n\nif optimizer_parameter_ids != trainable_parameter_ids:\n    raise RuntimeError(\n        \"Optimizer parameters do not exactly match \"\n        \"the trainable model parameters.\"\n    )\n\nprint(\n    \"Optimizer parameter coverage: PASS\"\n)\n\n# ----------------------------------------------------------------\n# SAVE PARAMETER SNAPSHOT\n#\n# Used only to verify that the optimizer step actually changes\n# trainable parameters and does not modify frozen parameters.\n# ----------------------------------------------------------------\n\nbefore_state = {\n    name: parameter.detach().clone()\n    for name, parameter in model.named_parameters()\n}\n\n# ----------------------------------------------------------------\n# LOAD ONE TRAINING BATCH\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"LOADING ONE TRAINING BATCH\")\nprint(\"=\" * 70)\n\nmodel.train()\n\nbatch = next(iter(train_loader))\n\nif not isinstance(batch, dict):\n    raise RuntimeError(\n        \"Expected dictionary batch.\"\n    )\n\nrequired_batch_keys = [\n    \"image\",\n    \"label\",\n    \"study_id\"\n]\n\nmissing_batch_keys = [\n    key\n    for key in required_batch_keys\n    if key not in batch\n]\n\nif missing_batch_keys:\n    raise RuntimeError(\n        \"Missing batch keys: \"\n        + \", \".join(missing_batch_keys)\n    )\n\nimages = batch[\"image\"].to(\n    device,\n    non_blocking=True\n)\n\nlabels = batch[\"label\"].to(\n    device,\n    non_blocking=True\n)\n\nprint(\n    f\"Images: {tuple(images.shape)}\"\n)\n\nprint(\n    f\"Labels: {tuple(labels.shape)}\"\n)\n\n# ----------------------------------------------------------------\n# INPUT VALIDATION\n# ----------------------------------------------------------------\n\nexpected_channels = 21\nexpected_targets = len(TARGETS)\n\nif images.ndim != 4:\n    raise RuntimeError(\n        f\"Expected 4D image tensor, got {images.ndim}D.\"\n    )\n\nif images.shape[1] != expected_channels:\n    raise RuntimeError(\n        f\"Expected {expected_channels} input channels, \"\n        f\"found {images.shape[1]}.\"\n    )\n\nif labels.ndim != 2:\n    raise RuntimeError(\n        f\"Expected 2D label tensor, got {labels.ndim}D.\"\n    )\n\nif labels.shape[1] != expected_targets:\n    raise RuntimeError(\n        f\"Expected {expected_targets} targets, \"\n        f\"found {labels.shape[1]}.\"\n    )\n\nif not torch.isfinite(images).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in input images.\"\n    )\n\nif not torch.isfinite(labels).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in labels.\"\n    )\n\nprint(\n    \"Input validation: PASS\"\n)\n\n# ----------------------------------------------------------------\n# LOSS\n#\n# Recreate the same class-weighted BCE formulation used in the\n# original training configuration.\n# ----------------------------------------------------------------\n\npositive_counts = labels.sum(\n    dim=0\n)\n\nnegative_counts = (\n    labels.shape[0]\n    - positive_counts\n)\n\nbatch_pos_weights = torch.where(\n    positive_counts > 0,\n    negative_counts / positive_counts,\n    torch.ones_like(positive_counts)\n)\n\nbatch_pos_weights = batch_pos_weights.to(\n    device=device,\n    dtype=torch.float32\n)\n\ncriterion_smoke = torch.nn.BCEWithLogitsLoss(\n    pos_weight=batch_pos_weights\n)\n\nprint(\n    \"Loss: BCEWithLogitsLoss\"\n)\n\n# ----------------------------------------------------------------\n# FORWARD PASS\n# ----------------------------------------------------------------\n\nfine_tune_optimizer.zero_grad(\n    set_to_none=True\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FORWARD PASS\")\nprint(\"=\" * 70)\n\nlogits = model(images)\n\nif logits.shape != labels.shape:\n    raise RuntimeError(\n        \"Output/label shape mismatch. \"\n        f\"Logits={tuple(logits.shape)}, \"\n        f\"Labels={tuple(labels.shape)}\"\n    )\n\nif not torch.isfinite(logits).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in logits.\"\n    )\n\nloss = criterion_smoke(\n    logits,\n    labels\n)\n\nif not torch.isfinite(loss):\n    raise RuntimeError(\n        \"Loss is NaN/Inf.\"\n    )\n\nprobabilities = torch.sigmoid(\n    logits\n)\n\nif not torch.isfinite(probabilities).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in probabilities.\"\n    )\n\nprint(\n    f\"Logit shape: {tuple(logits.shape)}\"\n)\n\nprint(\n    f\"Loss before update: {loss.item():.6f}\"\n)\n\nprint(\n    f\"Logit min: {logits.min().item():.6f}\"\n)\n\nprint(\n    f\"Logit max: {logits.max().item():.6f}\"\n)\n\nprint(\n    f\"Probability min: {probabilities.min().item():.6f}\"\n)\n\nprint(\n    f\"Probability max: {probabilities.max().item():.6f}\"\n)\n\nprint(\n    \"Forward pass: PASS\"\n)\n\n# ----------------------------------------------------------------\n# BACKWARD PASS\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"BACKWARD PASS\")\nprint(\"=\" * 70)\n\nloss.backward()\n\ngradient_values = []\n\ninvalid_gradients = []\n\nfor name, parameter in model.named_parameters():\n\n    if not parameter.requires_grad:\n        continue\n\n    if parameter.grad is None:\n        raise RuntimeError(\n            f\"Missing gradient for trainable parameter: {name}\"\n        )\n\n    if not torch.isfinite(\n        parameter.grad\n    ).all():\n        invalid_gradients.append(name)\n\n    gradient_values.append(\n        parameter.grad.detach().norm().item()\n    )\n\nif invalid_gradients:\n    raise RuntimeError(\n        \"NaN/Inf gradients detected:\\n\"\n        + \"\\n\".join(invalid_gradients)\n    )\n\ntotal_gradient_norm = torch.nn.utils.clip_grad_norm_(\n    model.parameters(),\n    max_norm=1.0\n)\n\nif not torch.isfinite(\n    torch.as_tensor(total_gradient_norm)\n):\n    raise RuntimeError(\n        \"Invalid gradient norm.\"\n    )\n\nprint(\n    f\"Gradient norm before clipping: \"\n    f\"{float(total_gradient_norm):.6f}\"\n)\n\nprint(\n    \"Gradient clipping: PASS\"\n)\n\nprint(\n    \"Backward pass: PASS\"\n)\n\n# ----------------------------------------------------------------\n# OPTIMIZER STEP\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"OPTIMIZER STEP\")\nprint(\"=\" * 70)\n\nfine_tune_optimizer.step()\n\nprint(\n    \"Optimizer step: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY PARAMETER CHANGES\n# ----------------------------------------------------------------\n\nchanged_trainable = []\nchanged_frozen = []\n\nfor name, parameter in model.named_parameters():\n\n    before = before_state[name]\n    after = parameter.detach()\n\n    changed = not torch.equal(\n        before,\n        after\n    )\n\n    if parameter.requires_grad:\n\n        if changed:\n            changed_trainable.append(name)\n\n    else:\n\n        if changed:\n            changed_frozen.append(name)\n\nif not changed_trainable:\n    raise RuntimeError(\n        \"No trainable parameters changed after optimizer step.\"\n    )\n\nif changed_frozen:\n    raise RuntimeError(\n        \"Frozen parameters changed unexpectedly:\\n\"\n        + \"\\n\".join(changed_frozen)\n    )\n\nprint(\n    f\"Trainable tensors changed: \"\n    f\"{len(changed_trainable)}\"\n)\n\nprint(\n    \"Frozen parameters unchanged: PASS\"\n)\n\nprint(\n    \"Optimizer parameter update: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY PARAMETERS REMAIN VALID\n# ----------------------------------------------------------------\n\ninvalid_parameters = []\n\nfor name, parameter in model.named_parameters():\n\n    if not torch.isfinite(\n        parameter.detach()\n    ).all():\n\n        invalid_parameters.append(name)\n\nif invalid_parameters:\n    raise RuntimeError(\n        \"NaN/Inf detected after optimizer step:\\n\"\n        + \"\\n\".join(invalid_parameters)\n    )\n\nprint(\n    \"Post-update parameter validity: PASS\"\n)\n\n# ----------------------------------------------------------------\n# DO NOT SAVE CHECKPOINT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 51 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    \"21-channel input: PASS\"\n)\n\nprint(\n    \"12-target output: PASS\"\n)\n\nprint(\n    \"Forward pass: PASS\"\n)\n\nprint(\n    \"Backward pass: PASS\"\n)\n\nprint(\n    \"Gradient clipping: PASS\"\n)\n\nprint(\n    \"Differential-learning-rate optimizer: PASS\"\n)\n\nprint(\n    \"Trainable parameter update: PASS\"\n)\n\nprint(\n    \"Frozen parameter protection: PASS\"\n)\n\nprint(\n    \"NaN/Inf validation: PASS\"\n)\n\nprint(\n    \"No checkpoint saved.\"\n)\n\nprint(\n    \"No existing fold checkpoint overwritten.\"\n)\n\nprint(\n    \"One training batch only.\"\n)\n\nprint(\n    \"\\nCELL 51 COMPLETE\"\n)\n\nprint(\n    \"Controlled fine-tuning smoke test completed.\"\n)\n\nprint(\n    \"Next step requires the complete Cell 51 output.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:36:38.252839Z","iopub.execute_input":"2026-08-11T10:36:38.253742Z","iopub.status.idle":"2026-08-11T10:36:39.767706Z","shell.execute_reply.started":"2026-08-11T10:36:38.253708Z","shell.execute_reply":"2026-08-11T10:36:39.766695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 52 - CONTROLLED FINE-TUNING FOLD 0 EXPERIMENT\n# ================================================================\n\nimport os\nimport copy\nimport torch\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 70)\nprint(\"CELL 52 - CONTROLLED FINE-TUNING FOLD 0 EXPERIMENT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\",\n    \"KneeStudyDataset\",\n    \"modeling_table\",\n    \"study_folds\"\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = next(model.parameters()).device\n\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# CONFIGURATION\n# ----------------------------------------------------------------\n\nFOLD = 0\nMAX_EPOCHS = 8\nEARLY_STOPPING_PATIENCE = 3\n\nCONV1_LR = 1e-4\nLAYER4_LR = 1e-5\nCLASSIFIER_LR = 1e-3\n\nWEIGHT_DECAY = 1e-4\nMAX_GRAD_NORM = 1.0\n\nIMAGE_SIZE = 224\nINPUT_CHANNELS = 21\nNUM_TARGETS = len(TARGETS)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"EXPERIMENT CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(f\"Fold: {FOLD}\")\nprint(f\"Maximum epochs: {MAX_EPOCHS}\")\nprint(f\"Early stopping patience: {EARLY_STOPPING_PATIENCE}\")\nprint(f\"Conv1 learning rate: {CONV1_LR}\")\nprint(f\"Layer4 learning rate: {LAYER4_LR}\")\nprint(f\"Classifier learning rate: {CLASSIFIER_LR}\")\nprint(f\"Weight decay: {WEIGHT_DECAY}\")\nprint(f\"Gradient clipping: {MAX_GRAD_NORM}\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# RESTORE MODEL TO ORIGINAL CELL-50 STRUCTURE\n# ----------------------------------------------------------------\n\nmodel.eval()\n\n# Freeze everything first\nfor parameter in model.parameters():\n    parameter.requires_grad = False\n\n# Conv1\nfor parameter in model.backbone.conv1.parameters():\n    parameter.requires_grad = True\n\n# Layer4\nfor parameter in model.backbone.layer4.parameters():\n    parameter.requires_grad = True\n\n# Classifier\nfor parameter in model.classifier.parameters():\n    parameter.requires_grad = True\n\n# ----------------------------------------------------------------\n# VERIFY TRAINABILITY\n# ----------------------------------------------------------------\n\nexpected_trainable_prefixes = [\n    \"backbone.conv1\",\n    \"backbone.layer4\",\n    \"classifier\"\n]\n\nunexpected_trainable = []\n\nfor name, parameter in model.named_parameters():\n\n    expected = any(\n        name.startswith(prefix)\n        for prefix in expected_trainable_prefixes\n    )\n\n    if parameter.requires_grad != expected:\n        unexpected_trainable.append(name)\n\nif unexpected_trainable:\n    raise RuntimeError(\n        \"Incorrect trainability configuration:\\n\"\n        + \"\\n\".join(unexpected_trainable)\n    )\n\ntrainable_count = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"TRAINABILITY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Trainable parameters: {trainable_count:,}\"\n)\n\nif trainable_count != 8_593_996:\n    raise RuntimeError(\n        f\"Unexpected trainable parameter count: {trainable_count}\"\n    )\n\nprint(\"Trainability configuration: PASS\")\n\n# ----------------------------------------------------------------\n# BUILD FOLD-0 TABLES\n# ----------------------------------------------------------------\n\nif \"fold\" not in study_folds.columns:\n    raise RuntimeError(\n        \"Fold table does not contain 'fold' column.\"\n    )\n\nfold0_ids = set(\n    study_folds.loc[\n        study_folds[\"fold\"] == FOLD,\n        \"StudyInstanceUID\"\n    ].astype(str)\n)\n\nif len(fold0_ids) != 20:\n    raise RuntimeError(\n        f\"Expected 20 Fold-0 validation studies, \"\n        f\"found {len(fold0_ids)}.\"\n    )\n\nall_ids = set(\n    modeling_table[\"StudyInstanceUID\"].astype(str)\n)\n\ntrain_ids = all_ids - fold0_ids\n\nif len(train_ids) != 38:\n    raise RuntimeError(\n        f\"Expected 38 Fold-0 training studies, \"\n        f\"found {len(train_ids)}.\"\n    )\n\ntrain_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].astype(str).isin(train_ids)\n].copy()\n\nval_table = modeling_table[\n    modeling_table[\"StudyInstanceUID\"].astype(str).isin(fold0_ids)\n].copy()\n\nif set(train_table[\"StudyInstanceUID\"].astype(str)) & \\\n   set(val_table[\"StudyInstanceUID\"].astype(str)):\n    raise RuntimeError(\n        \"Train/validation overlap detected.\"\n    )\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 0 DATA\")\nprint(\"=\" * 70)\n\nprint(f\"Training studies: {len(train_table)}\")\nprint(f\"Validation studies: {len(val_table)}\")\nprint(\"Train/validation overlap: 0\")\n\n# ----------------------------------------------------------------\n# BUILD DATASETS USING THE ESTABLISHED PIPELINE\n# ----------------------------------------------------------------\n\ntrain_dataset_ft = KneeStudyDataset(\n    train_table,\n    targets=TARGETS\n)\n\nval_dataset_ft = KneeStudyDataset(\n    val_table,\n    targets=TARGETS\n)\n\n# ----------------------------------------------------------------\n# DATALOADERS\n# ----------------------------------------------------------------\n\ntrain_loader_ft = torch.utils.data.DataLoader(\n    train_dataset_ft,\n    batch_size=2,\n    shuffle=True,\n    num_workers=0,\n    pin_memory=False\n)\n\nval_loader_ft = torch.utils.data.DataLoader(\n    val_dataset_ft,\n    batch_size=2,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False\n)\n\nprint(f\"Training batches: {len(train_loader_ft)}\")\nprint(f\"Validation batches: {len(val_loader_ft)}\")\n\n# ----------------------------------------------------------------\n# COMPUTE TRAINING CLASS WEIGHTS\n# ----------------------------------------------------------------\n\ntrain_labels = train_table[TARGETS].astype(float)\n\npositive_counts = train_labels.sum(axis=0)\nnegative_counts = len(train_labels) - positive_counts\n\npos_weight = np.where(\n    positive_counts > 0,\n    negative_counts / positive_counts,\n    1.0\n)\n\npos_weight_tensor = torch.tensor(\n    pos_weight,\n    dtype=torch.float32,\n    device=device\n)\n\ncriterion = torch.nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_tensor\n)\n\nprint(\"\\nClass-weighted BCE: PASS\")\n\n# ----------------------------------------------------------------\n# DIFFERENTIAL LEARNING-RATE OPTIMIZER\n# ----------------------------------------------------------------\n\noptimizer = torch.optim.AdamW(\n    [\n        {\n            \"params\": model.backbone.conv1.parameters(),\n            \"lr\": CONV1_LR\n        },\n        {\n            \"params\": model.backbone.layer4.parameters(),\n            \"lr\": LAYER4_LR\n        },\n        {\n            \"params\": model.classifier.parameters(),\n            \"lr\": CLASSIFIER_LR\n        }\n    ],\n    weight_decay=WEIGHT_DECAY\n)\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"min\",\n    factor=0.5,\n    patience=2\n)\n\nprint(\"Differential-learning-rate optimizer: PASS\")\n\n# ----------------------------------------------------------------\n# METRIC FUNCTION\n# ----------------------------------------------------------------\n\ndef compute_epoch_metrics(\n    probabilities,\n    labels,\n    threshold=0.50\n):\n\n    predictions = (\n        probabilities >= threshold\n    ).astype(np.int32)\n\n    f1_values = []\n\n    for target_index in range(\n        labels.shape[1]\n    ):\n\n        y_true = labels[:, target_index]\n        y_pred = predictions[:, target_index]\n\n        tp = np.sum(\n            (y_true == 1) & (y_pred == 1)\n        )\n\n        fp = np.sum(\n            (y_true == 0) & (y_pred == 1)\n        )\n\n        fn = np.sum(\n            (y_true == 1) & (y_pred == 0)\n        )\n\n        denominator = (\n            2 * tp + fp + fn\n        )\n\n        if denominator == 0:\n            f1 = 0.0\n        else:\n            f1 = (\n                2 * tp\n                / denominator\n            )\n\n        f1_values.append(float(f1))\n\n    return float(np.mean(f1_values))\n\n# ----------------------------------------------------------------\n# VALIDATION FUNCTION\n# ----------------------------------------------------------------\n\ndef evaluate_model():\n\n    model.eval()\n\n    total_loss = 0.0\n    all_probabilities = []\n    all_labels = []\n\n    with torch.no_grad():\n\n        for batch in val_loader_ft:\n\n            images = batch[\"image\"].to(\n                device,\n                non_blocking=True\n            )\n\n            labels = batch[\"label\"].to(\n                device,\n                non_blocking=True\n            )\n\n            logits = model(images)\n\n            loss = criterion(\n                logits,\n                labels\n            )\n\n            probabilities = torch.sigmoid(\n                logits\n            )\n\n            total_loss += (\n                loss.item()\n                * images.size(0)\n            )\n\n            all_probabilities.append(\n                probabilities.cpu().numpy()\n            )\n\n            all_labels.append(\n                labels.cpu().numpy()\n            )\n\n    probabilities = np.concatenate(\n        all_probabilities,\n        axis=0\n    )\n\n    labels = np.concatenate(\n        all_labels,\n        axis=0\n    )\n\n    average_loss = (\n        total_loss\n        / len(val_dataset_ft)\n    )\n\n    mean_f1 = compute_epoch_metrics(\n        probabilities,\n        labels\n    )\n\n    return (\n        average_loss,\n        mean_f1\n    )\n\n# ----------------------------------------------------------------\n# TRAINING\n# ----------------------------------------------------------------\n\nhistory = []\n\nbest_val_loss = float(\"inf\")\nbest_epoch = 0\nepochs_without_improvement = 0\n\nbest_state = None\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING CONTROLLED FOLD 0 FINE-TUNING\")\nprint(\"=\" * 70)\n\nfor epoch in range(\n    1,\n    MAX_EPOCHS + 1\n):\n\n    model.train()\n\n    running_loss = 0.0\n    sample_count = 0\n\n    for batch in train_loader_ft:\n\n        images = batch[\"image\"].to(\n            device,\n            non_blocking=True\n        )\n\n        labels = batch[\"label\"].to(\n            device,\n            non_blocking=True\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        logits = model(images)\n\n        loss = criterion(\n            logits,\n            labels\n        )\n\n        if not torch.isfinite(loss):\n            raise RuntimeError(\n                f\"Non-finite training loss at epoch {epoch}.\"\n            )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=MAX_GRAD_NORM\n        )\n\n        optimizer.step()\n\n        batch_size = images.size(0)\n\n        running_loss += (\n            loss.item()\n            * batch_size\n        )\n\n        sample_count += batch_size\n\n    train_loss = (\n        running_loss\n        / sample_count\n    )\n\n    val_loss, mean_f1 = evaluate_model()\n\n    scheduler.step(val_loss)\n\n    current_lr = [\n        group[\"lr\"]\n        for group in optimizer.param_groups\n    ]\n\n    history.append(\n        {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"mean_f1\": mean_f1,\n            \"conv1_lr\": current_lr[0],\n            \"layer4_lr\": current_lr[1],\n            \"classifier_lr\": current_lr[2]\n        }\n    )\n\n    if val_loss < best_val_loss:\n\n        best_val_loss = val_loss\n        best_epoch = epoch\n        epochs_without_improvement = 0\n\n        best_state = copy.deepcopy(\n            model.state_dict()\n        )\n\n    else:\n\n        epochs_without_improvement += 1\n\n    print(\n        f\"\\nEpoch {epoch:02d}/{MAX_EPOCHS}\"\n    )\n\n    print(\n        f\"Train Loss: {train_loss:.5f}\"\n    )\n\n    print(\n        f\"Val Loss:   {val_loss:.5f}\"\n    )\n\n    print(\n        f\"Mean F1:    {mean_f1:.4f}\"\n    )\n\n    print(\n        \"Learning Rates: \"\n        f\"Conv1={current_lr[0]:.6f} | \"\n        f\"Layer4={current_lr[1]:.6f} | \"\n        f\"Classifier={current_lr[2]:.6f}\"\n    )\n\n    print(\n        f\"Best Epoch: {best_epoch}\"\n    )\n\n    print(\n        f\"Best Val Loss: {best_val_loss:.5f}\"\n    )\n\n    if epochs_without_improvement >= \\\n       EARLY_STOPPING_PATIENCE:\n\n        print(\n            \"\\nEarly stopping triggered.\"\n        )\n\n        break\n\n# ----------------------------------------------------------------\n# RESTORE BEST EXPERIMENT STATE\n# ----------------------------------------------------------------\n\nif best_state is None:\n    raise RuntimeError(\n        \"No valid fine-tuning checkpoint state was produced.\"\n    )\n\nmodel.load_state_dict(\n    best_state\n)\n\nmodel.eval()\n\n# ----------------------------------------------------------------\n# SAVE ONLY EXPERIMENTAL OUTPUT\n# ----------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nos.makedirs(\n    audit_dir,\n    exist_ok=True\n)\n\ncheckpoint_path = os.path.join(\n    audit_dir,\n    \"cell52_fold0_finetuned_experiment.pt\"\n)\n\nhistory_path = os.path.join(\n    audit_dir,\n    \"cell52_fold0_finetuned_history.csv\"\n)\n\ntorch.save(\n    model.state_dict(),\n    checkpoint_path\n)\n\npd.DataFrame(history).to_csv(\n    history_path,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL OUTPUT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 52 - EXPERIMENT COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Best epoch: {best_epoch}\"\n)\n\nprint(\n    f\"Best validation loss: \"\n    f\"{best_val_loss:.6f}\"\n)\n\nprint(\n    f\"Final validation loss: \"\n    f\"{best_val_loss:.6f}\"\n)\n\nprint(\n    f\"Experimental checkpoint saved: \"\n    f\"{checkpoint_path}\"\n)\n\nprint(\n    f\"History saved: \"\n    f\"{history_path}\"\n)\n\nprint(\n    \"\\nIMPORTANT:\"\n)\n\nprint(\n    \"Original Cell-37 Fold-0 checkpoint was NOT overwritten.\"\n)\n\nprint(\n    \"This is a controlled fine-tuning experiment only.\"\n)\n\nprint(\n    \"\\nCELL 52 COMPLETE\"\n)\n\nprint(\n    \"Send the complete Cell 52 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:39:15.588406Z","iopub.execute_input":"2026-08-11T10:39:15.58873Z","iopub.status.idle":"2026-08-11T10:41:50.887839Z","shell.execute_reply.started":"2026-08-11T10:39:15.588705Z","shell.execute_reply":"2026-08-11T10:41:50.886722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 53 - FINE-TUNED FOLD 0 DIAGNOSTIC AUDIT\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score\n\nprint(\"=\" * 70)\nprint(\"CELL 53 - FINE-TUNED FOLD 0 DIAGNOSTIC AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"val_loader_ft\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nprint(\"\\nRequired notebook objects: PASS\")\n\n# ----------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------\n\ndevice = next(model.parameters()).device\n\nprint(\n    f\"Device: {device}\"\n)\n\n# ----------------------------------------------------------------\n# CHECK EXPERIMENTAL CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell52_fold0_finetuned_experiment.pt\"\n)\n\nif not os.path.exists(checkpoint_path):\n    raise RuntimeError(\n        \"Fine-tuned Fold-0 checkpoint not found:\\n\"\n        + checkpoint_path\n    )\n\nprint(\n    \"Fine-tuned Fold-0 checkpoint: EXISTS\"\n)\n\n# ----------------------------------------------------------------\n# LOAD EXPERIMENTAL CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint_state = torch.load(\n    checkpoint_path,\n    map_location=device\n)\n\nmodel.load_state_dict(\n    checkpoint_state\n)\n\nmodel.eval()\n\nprint(\n    \"Fine-tuned Fold-0 checkpoint: LOADED\"\n)\n\nprint(\n    \"Model mode: evaluation\"\n)\n\n# ----------------------------------------------------------------\n# COLLECT VALIDATION PREDICTIONS\n# ----------------------------------------------------------------\n\nall_probabilities = []\nall_labels = []\nall_study_ids = []\n\nwith torch.no_grad():\n\n    for batch in val_loader_ft:\n\n        images = batch[\"image\"].to(\n            device,\n            non_blocking=True\n        )\n\n        labels = batch[\"label\"].to(\n            device,\n            non_blocking=True\n        )\n\n        logits = model(images)\n\n        if logits.shape != labels.shape:\n            raise RuntimeError(\n                \"Logit/label shape mismatch: \"\n                f\"{tuple(logits.shape)} vs \"\n                f\"{tuple(labels.shape)}\"\n            )\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n        if not torch.isfinite(\n            probabilities\n        ).all():\n            raise RuntimeError(\n                \"NaN/Inf detected in probabilities.\"\n            )\n\n        all_probabilities.append(\n            probabilities.cpu().numpy()\n        )\n\n        all_labels.append(\n            labels.cpu().numpy()\n        )\n\n        all_study_ids.extend(\n            batch[\"study_id\"]\n        )\n\nprobabilities = np.concatenate(\n    all_probabilities,\n    axis=0\n)\n\nlabels = np.concatenate(\n    all_labels,\n    axis=0\n)\n\n# ----------------------------------------------------------------\n# BASIC VALIDATION\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"PREDICTION VALIDATION\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Probability shape: {probabilities.shape}\"\n)\n\nprint(\n    f\"Label shape: {labels.shape}\"\n)\n\nif probabilities.shape != labels.shape:\n    raise RuntimeError(\n        \"Prediction/label shape mismatch.\"\n    )\n\nif probabilities.shape[1] != len(TARGETS):\n    raise RuntimeError(\n        \"Unexpected target dimension.\"\n    )\n\nif len(all_study_ids) != probabilities.shape[0]:\n    raise RuntimeError(\n        \"Study ID count does not match predictions.\"\n    )\n\nif len(set(all_study_ids)) != len(all_study_ids):\n    raise RuntimeError(\n        \"Duplicate validation study predictions detected.\"\n    )\n\nprint(\n    \"Probability validity: PASS\"\n)\n\nprint(\n    \"Shape consistency: PASS\"\n)\n\nprint(\n    \"One prediction per validation study: PASS\"\n)\n\n# ----------------------------------------------------------------\n# ROC-AUC + PROBABILITY SEPARATION\n# ----------------------------------------------------------------\n\nrows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    positive_mask = (\n        y_true == 1\n    )\n\n    negative_mask = (\n        y_true == 0\n    )\n\n    positive_count = int(\n        positive_mask.sum()\n    )\n\n    negative_count = int(\n        negative_mask.sum()\n    )\n\n    positive_mean = (\n        float(y_prob[positive_mask].mean())\n        if positive_count > 0\n        else np.nan\n    )\n\n    negative_mean = (\n        float(y_prob[negative_mask].mean())\n        if negative_count > 0\n        else np.nan\n    )\n\n    positive_median = (\n        float(np.median(y_prob[positive_mask]))\n        if positive_count > 0\n        else np.nan\n    )\n\n    negative_median = (\n        float(np.median(y_prob[negative_mask]))\n        if negative_count > 0\n        else np.nan\n    )\n\n    separation = (\n        positive_mean\n        - negative_mean\n    )\n\n    if (\n        positive_count > 0\n        and negative_count > 0\n    ):\n        roc_auc = roc_auc_score(\n            y_true,\n            y_prob\n        )\n    else:\n        roc_auc = np.nan\n\n    rows.append(\n        {\n            \"target\": target,\n            \"positive_count\": positive_count,\n            \"negative_count\": negative_count,\n            \"positive_mean\": positive_mean,\n            \"negative_mean\": negative_mean,\n            \"positive_median\": positive_median,\n            \"negative_median\": negative_median,\n            \"probability_separation\": separation,\n            \"roc_auc\": roc_auc\n        }\n    )\n\nseparation_df = pd.DataFrame(rows)\n\n# ----------------------------------------------------------------\n# THRESHOLD AUDIT\n# ----------------------------------------------------------------\n\nthresholds = np.arange(\n    0.10,\n    0.901,\n    0.025\n)\n\nthreshold_rows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    best_threshold = 0.50\n    best_f1 = 0.0\n    f1_at_050 = 0.0\n\n    for threshold in thresholds:\n\n        y_pred = (\n            y_prob >= threshold\n        ).astype(np.int32)\n\n        tp = np.sum(\n            (y_true == 1)\n            & (y_pred == 1)\n        )\n\n        fp = np.sum(\n            (y_true == 0)\n            & (y_pred == 1)\n        )\n\n        fn = np.sum(\n            (y_true == 1)\n            & (y_pred == 0)\n        )\n\n        denominator = (\n            2 * tp\n            + fp\n            + fn\n        )\n\n        if denominator == 0:\n            f1 = 0.0\n        else:\n            f1 = (\n                2 * tp\n                / denominator\n            )\n\n        if abs(threshold - 0.50) < 1e-8:\n            f1_at_050 = float(f1)\n\n        if f1 > best_f1:\n            best_f1 = float(f1)\n            best_threshold = float(threshold)\n\n    threshold_rows.append(\n        {\n            \"target\": target,\n            \"f1_at_0.50\": f1_at_050,\n            \"best_threshold\": best_threshold,\n            \"best_f1\": best_f1,\n            \"f1_improvement\": (\n                best_f1 - f1_at_050\n            )\n        }\n    )\n\nthreshold_df = pd.DataFrame(\n    threshold_rows\n)\n\n# ----------------------------------------------------------------\n# MERGE DIAGNOSTICS\n# ----------------------------------------------------------------\n\ndiagnostics_df = separation_df.merge(\n    threshold_df,\n    on=\"target\",\n    how=\"inner\"\n)\n\n# ----------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------\n\nmean_roc_auc = float(\n    diagnostics_df[\"roc_auc\"].mean()\n)\n\nmean_separation = float(\n    diagnostics_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    diagnostics_df[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nmean_best_f1 = float(\n    diagnostics_df[\n        \"best_f1\"\n    ].mean()\n)\n\nf1_improvement = (\n    mean_best_f1\n    - mean_f1_050\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINE-TUNED FOLD 0 DIAGNOSTICS\")\nprint(\"=\" * 70)\n\nprint(\n    diagnostics_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Mean ROC-AUC: {mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50: {mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{f1_improvement:.4f}\"\n)\n\nhigh_auc_targets = diagnostics_df.loc[\n    diagnostics_df[\"roc_auc\"] >= 0.65,\n    \"target\"\n].tolist()\n\nnegative_separation_targets = diagnostics_df.loc[\n    diagnostics_df[\n        \"probability_separation\"\n    ] < 0,\n    \"target\"\n].tolist()\n\nlarge_f1_gain_targets = diagnostics_df.loc[\n    diagnostics_df[\n        \"f1_improvement\"\n    ] > 0.05,\n    \"target\"\n].tolist()\n\nprint(\n    f\"\\nTargets with ROC-AUC >= 0.65: \"\n    f\"{len(high_auc_targets)}\"\n)\n\nprint(\n    high_auc_targets\n)\n\nprint(\n    f\"\\nTargets with negative probability separation: \"\n    f\"{len(negative_separation_targets)}\"\n)\n\nprint(\n    negative_separation_targets\n)\n\nprint(\n    f\"\\nTargets with >0.05 F1 improvement: \"\n    f\"{len(large_f1_gain_targets)}\"\n)\n\nprint(\n    large_f1_gain_targets\n)\n\n# ----------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------\n\naudit_dir = (\n    \"/kaggle/working/rsna_knee_audit\"\n)\n\ndiagnostics_path = os.path.join(\n    audit_dir,\n    \"cell53_fold0_finetuned_diagnostics.csv\"\n)\n\nthreshold_path = os.path.join(\n    audit_dir,\n    \"cell53_fold0_finetuned_threshold_audit.csv\"\n)\n\ndiagnostics_df.to_csv(\n    diagnostics_path,\n    index=False\n)\n\nthreshold_df.to_csv(\n    threshold_path,\n    index=False\n)\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 53 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Saved: {diagnostics_path}\"\n)\n\nprint(\n    f\"Saved: {threshold_path}\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No original checkpoint modified.\"\n)\n\nprint(\n    \"No test data used.\"\n)\n\nprint(\n    \"Send the complete Cell 53 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:44:35.222842Z","iopub.execute_input":"2026-08-11T10:44:35.223327Z","iopub.status.idle":"2026-08-11T10:44:44.741429Z","shell.execute_reply.started":"2026-08-11T10:44:35.223296Z","shell.execute_reply":"2026-08-11T10:44:44.740432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 54R - CONTROLLED FINE-TUNING FOLD 1\n# CORRECTED CONTINUATION AFTER MODEL-RESET FAILURE\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\n\nprint(\"=\" * 70)\nprint(\"CELL 54R - CONTROLLED FINE-TUNING FOLD 1\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# VERIFY OBJECTS CREATED SUCCESSFULLY BY CELL 54\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"train_loader_ft\",\n    \"val_loader_ft\",\n    \"train_table\",\n    \"val_table\",\n    \"TARGETS\",\n    \"device\"\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# PATHS\n# ----------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nbaseline_checkpoint = os.path.join(\n    audit_dir,\n    \"cell40_fold1_best_model.pt\"\n)\n\nexperimental_checkpoint = os.path.join(\n    audit_dir,\n    \"cell54_fold1_finetuned_experiment.pt\"\n)\n\nhistory_path = os.path.join(\n    audit_dir,\n    \"cell54_fold1_finetuned_history.csv\"\n)\n\nmetrics_path = os.path.join(\n    audit_dir,\n    \"cell54_fold1_finetuned_metrics.csv\"\n)\n\nif not os.path.exists(baseline_checkpoint):\n    raise RuntimeError(\n        \"Original Cell-40 Fold-1 checkpoint not found: \"\n        + baseline_checkpoint\n    )\n\nprint(\n    \"Original Fold-1 baseline checkpoint: EXISTS\"\n)\n\n# ----------------------------------------------------------------\n# LOAD ORIGINAL FOLD-1 BASELINE MODEL\n# ----------------------------------------------------------------\n# IMPORTANT:\n# Cell 40 is the baseline Fold-1 model.\n# Cell 54 fine-tuning must start from that checkpoint.\n# Cell 40 checkpoint itself is never modified.\n\nbaseline_state = torch.load(\n    baseline_checkpoint,\n    map_location=device\n)\n\n# Support both a plain state_dict and a checkpoint dictionary,\n# without assuming a checkpoint structure that was not established.\n\nif isinstance(baseline_state, dict):\n\n    if \"state_dict\" in baseline_state:\n        baseline_state_dict = baseline_state[\"state_dict\"]\n\n    elif \"model_state_dict\" in baseline_state:\n        baseline_state_dict = baseline_state[\"model_state_dict\"]\n\n    else:\n        baseline_state_dict = baseline_state\n\nelse:\n    raise RuntimeError(\n        \"Unsupported Fold-1 checkpoint format.\"\n    )\n\nmodel.load_state_dict(\n    baseline_state_dict,\n    strict=True\n)\n\nmodel = model.to(device)\n\nprint(\n    \"Original Cell-40 Fold-1 checkpoint loaded into existing model: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY MODEL ARCHITECTURE\n# ----------------------------------------------------------------\n\nif not hasattr(model, \"backbone\"):\n    raise RuntimeError(\n        \"Existing model does not contain 'backbone'.\"\n    )\n\nif not hasattr(model.backbone, \"conv1\"):\n    raise RuntimeError(\n        \"Existing model backbone does not contain 'conv1'.\"\n    )\n\nif not hasattr(model.backbone, \"layer1\"):\n    raise RuntimeError(\n        \"Existing model backbone does not contain 'layer1'.\"\n    )\n\nif not hasattr(model.backbone, \"layer2\"):\n    raise RuntimeError(\n        \"Existing model backbone does not contain 'layer2'.\"\n    )\n\nif not hasattr(model.backbone, \"layer3\"):\n    raise RuntimeError(\n        \"Existing model backbone does not contain 'layer3'.\"\n    )\n\nif not hasattr(model.backbone, \"layer4\"):\n    raise RuntimeError(\n        \"Existing model backbone does not contain 'layer4'.\"\n    )\n\nif not hasattr(model, \"classifier\"):\n    raise RuntimeError(\n        \"Existing model does not contain 'classifier'.\"\n    )\n\nprint(\"Existing model architecture: PASS\")\n\n# ----------------------------------------------------------------\n# VERIFY 21-CHANNEL INPUT AND 12-CLASS OUTPUT\n# ----------------------------------------------------------------\n\nsmoke_batch = next(iter(train_loader_ft))\n\nimages = smoke_batch[\"image\"].to(device)\nlabels = smoke_batch[\"label\"].to(device)\n\nif tuple(images.shape) != (2, 21, 224, 224):\n    raise RuntimeError(\n        \"Unexpected input shape: \"\n        + str(tuple(images.shape))\n    )\n\nif tuple(labels.shape) != (2, 12):\n    raise RuntimeError(\n        \"Unexpected label shape: \"\n        + str(tuple(labels.shape))\n    )\n\nmodel.eval()\n\nwith torch.no_grad():\n    smoke_logits = model(images)\n\nif tuple(smoke_logits.shape) != (2, 12):\n    raise RuntimeError(\n        \"Unexpected model output shape: \"\n        + str(tuple(smoke_logits.shape))\n    )\n\nif not torch.isfinite(smoke_logits).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in model output.\"\n    )\n\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Forward-pass validation: PASS\")\n\n# ----------------------------------------------------------------\n# CONTROLLED TRAINABILITY\n# ----------------------------------------------------------------\n\nfor parameter in model.parameters():\n    parameter.requires_grad = False\n\n# Conv1 trainable\nfor parameter in model.backbone.conv1.parameters():\n    parameter.requires_grad = True\n\n# Layer4 trainable\nfor parameter in model.backbone.layer4.parameters():\n    parameter.requires_grad = True\n\n# Classifier trainable\nfor parameter in model.classifier.parameters():\n    parameter.requires_grad = True\n\n# ----------------------------------------------------------------\n# TRAINABILITY VERIFICATION\n# ----------------------------------------------------------------\n\nlayer1_trainable = any(\n    parameter.requires_grad\n    for parameter in model.backbone.layer1.parameters()\n)\n\nlayer2_trainable = any(\n    parameter.requires_grad\n    for parameter in model.backbone.layer2.parameters()\n)\n\nlayer3_trainable = any(\n    parameter.requires_grad\n    for parameter in model.backbone.layer3.parameters()\n)\n\nconv1_trainable = any(\n    parameter.requires_grad\n    for parameter in model.backbone.conv1.parameters()\n)\n\nlayer4_trainable = any(\n    parameter.requires_grad\n    for parameter in model.backbone.layer4.parameters()\n)\n\nclassifier_trainable = any(\n    parameter.requires_grad\n    for parameter in model.classifier.parameters()\n)\n\nif layer1_trainable:\n    raise RuntimeError(\n        \"Layer1 is unexpectedly trainable.\"\n    )\n\nif layer2_trainable:\n    raise RuntimeError(\n        \"Layer2 is unexpectedly trainable.\"\n    )\n\nif layer3_trainable:\n    raise RuntimeError(\n        \"Layer3 is unexpectedly trainable.\"\n    )\n\nif not conv1_trainable:\n    raise RuntimeError(\n        \"Conv1 is not trainable.\"\n    )\n\nif not layer4_trainable:\n    raise RuntimeError(\n        \"Layer4 is not trainable.\"\n    )\n\nif not classifier_trainable:\n    raise RuntimeError(\n        \"Classifier is not trainable.\"\n    )\n\nprint(\"Controlled trainability: PASS\")\nprint(\"Conv1 trainable: YES\")\nprint(\"Layer1 frozen: YES\")\nprint(\"Layer2 frozen: YES\")\nprint(\"Layer3 frozen: YES\")\nprint(\"Layer4 trainable: YES\")\nprint(\"Classifier trainable: YES\")\n\n# ----------------------------------------------------------------\n# TRAINABLE PARAMETER COUNT\n# ----------------------------------------------------------------\n\ntrainable_parameter_count = sum(\n    parameter.numel()\n    for parameter in model.parameters()\n    if parameter.requires_grad\n)\n\nif trainable_parameter_count != 8_593_996:\n    raise RuntimeError(\n        \"Unexpected trainable parameter count: \"\n        + str(trainable_parameter_count)\n    )\n\nprint(\n    f\"Trainable parameters: {trainable_parameter_count:,}\"\n)\nprint(\"Trainable parameter count: PASS\")\n\n# ----------------------------------------------------------------\n# FOLD 1 CLASS WEIGHTS\n# ----------------------------------------------------------------\n\npos_weights = []\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FOLD 1 CLASS WEIGHTS\")\nprint(\"=\" * 70)\n\nfor target in TARGETS:\n\n    positive = int(\n        train_table[target].sum()\n    )\n\n    negative = (\n        len(train_table) - positive\n    )\n\n    if positive <= 0:\n        raise RuntimeError(\n            f\"No positive training samples for {target}.\"\n        )\n\n    pos_weight = (\n        negative / positive\n    )\n\n    pos_weights.append(\n        pos_weight\n    )\n\n    print(\n        f\"{target:<20} \"\n        f\"positive={positive:2d} \"\n        f\"negative={negative:2d} \"\n        f\"pos_weight={pos_weight:.4f}\"\n    )\n\npos_weight_tensor = torch.tensor(\n    pos_weights,\n    dtype=torch.float32,\n    device=device\n)\n\ncriterion = nn.BCEWithLogitsLoss(\n    pos_weight=pos_weight_tensor\n)\n\nprint(\"Class-weighted BCE: PASS\")\n\n# ----------------------------------------------------------------\n# DIFFERENTIAL LEARNING-RATE OPTIMIZER\n# ----------------------------------------------------------------\n\noptimizer = torch.optim.AdamW(\n    [\n        {\n            \"params\": model.backbone.conv1.parameters(),\n            \"lr\": 1e-4\n        },\n        {\n            \"params\": model.backbone.layer4.parameters(),\n            \"lr\": 1e-5\n        },\n        {\n            \"params\": model.classifier.parameters(),\n            \"lr\": 1e-3\n        }\n    ],\n    weight_decay=1e-4\n)\n\nprint(\"Differential-learning-rate optimizer: PASS\")\n\n# ----------------------------------------------------------------\n# EXPERIMENT CONFIGURATION\n# ----------------------------------------------------------------\n\nmax_epochs = 8\npatience = 3\ngradient_clip_norm = 1.0\n\nbest_val_loss = float(\"inf\")\nbest_epoch = 0\nepochs_without_improvement = 0\n\nhistory = []\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"EXPERIMENT CONFIGURATION\")\nprint(\"=\" * 70)\n\nprint(\"Fold: 1\")\nprint(f\"Maximum epochs: {max_epochs}\")\nprint(f\"Early stopping patience: {patience}\")\nprint(\"Conv1 learning rate: 0.0001\")\nprint(\"Layer4 learning rate: 0.00001\")\nprint(\"Classifier learning rate: 0.001\")\nprint(\"Weight decay: 0.0001\")\nprint(\"Gradient clipping: 1.0\")\nprint(f\"Device: {device}\")\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"STARTING CONTROLLED FOLD 1 FINE-TUNING\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# TRAINING LOOP\n# ----------------------------------------------------------------\n\nfor epoch in range(1, max_epochs + 1):\n\n    model.train()\n\n    train_loss_sum = 0.0\n    train_batch_count = 0\n\n    for batch in train_loader_ft:\n\n        batch_images = batch[\"image\"].to(\n            device\n        )\n\n        batch_labels = batch[\"label\"].to(\n            device\n        )\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        logits = model(\n            batch_images\n        )\n\n        loss = criterion(\n            logits,\n            batch_labels\n        )\n\n        if not torch.isfinite(loss):\n            raise RuntimeError(\n                \"NaN/Inf training loss detected.\"\n            )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=gradient_clip_norm\n        )\n\n        optimizer.step()\n\n        train_loss_sum += loss.item()\n        train_batch_count += 1\n\n    train_loss = (\n        train_loss_sum\n        / train_batch_count\n    )\n\n    # ------------------------------------------------------------\n    # VALIDATION\n    # ------------------------------------------------------------\n\n    model.eval()\n\n    val_loss_sum = 0.0\n    val_batch_count = 0\n\n    all_val_probabilities = []\n    all_val_labels = []\n\n    with torch.no_grad():\n\n        for batch in val_loader_ft:\n\n            batch_images = batch[\"image\"].to(\n                device\n            )\n\n            batch_labels = batch[\"label\"].to(\n                device\n            )\n\n            logits = model(\n                batch_images\n            )\n\n            loss = criterion(\n                logits,\n                batch_labels\n            )\n\n            probabilities = torch.sigmoid(\n                logits\n            )\n\n            if not torch.isfinite(loss):\n                raise RuntimeError(\n                    \"NaN/Inf validation loss detected.\"\n                )\n\n            if not torch.isfinite(\n                probabilities\n            ).all():\n                raise RuntimeError(\n                    \"NaN/Inf validation probabilities detected.\"\n                )\n\n            val_loss_sum += loss.item()\n            val_batch_count += 1\n\n            all_val_probabilities.append(\n                probabilities.cpu()\n            )\n\n            all_val_labels.append(\n                batch_labels.cpu()\n            )\n\n    val_loss = (\n        val_loss_sum\n        / val_batch_count\n    )\n\n    val_probabilities = torch.cat(\n        all_val_probabilities,\n        dim=0\n    ).numpy()\n\n    val_labels_np = torch.cat(\n        all_val_labels,\n        dim=0\n    ).numpy()\n\n    if val_probabilities.shape != (\n        len(val_table),\n        len(TARGETS)\n    ):\n        raise RuntimeError(\n            \"Unexpected validation probability shape: \"\n            + str(val_probabilities.shape)\n        )\n\n    # ------------------------------------------------------------\n    # F1 @ 0.50\n    # ------------------------------------------------------------\n\n    f1_values = []\n\n    for target_index in range(\n        len(TARGETS)\n    ):\n\n        y_true = val_labels_np[\n            :, target_index\n        ]\n\n        y_pred = (\n            val_probabilities[\n                :, target_index\n            ] >= 0.50\n        ).astype(np.int32)\n\n        tp = np.sum(\n            (y_true == 1)\n            & (y_pred == 1)\n        )\n\n        fp = np.sum(\n            (y_true == 0)\n            & (y_pred == 1)\n        )\n\n        fn = np.sum(\n            (y_true == 1)\n            & (y_pred == 0)\n        )\n\n        denominator = (\n            2 * tp + fp + fn\n        )\n\n        if denominator == 0:\n            f1 = 0.0\n        else:\n            f1 = (\n                2 * tp\n                / denominator\n            )\n\n        f1_values.append(\n            f1\n        )\n\n    mean_f1 = float(\n        np.mean(f1_values)\n    )\n\n    # ------------------------------------------------------------\n    # REDUCE LR ON PLATEAU\n    # ------------------------------------------------------------\n\n    if val_loss < best_val_loss - 1e-6:\n\n        best_val_loss = val_loss\n        best_epoch = epoch\n        epochs_without_improvement = 0\n\n        torch.save(\n            model.state_dict(),\n            experimental_checkpoint\n        )\n\n    else:\n\n        epochs_without_improvement += 1\n\n        if epochs_without_improvement >= 2:\n\n            for group in optimizer.param_groups:\n                group[\"lr\"] *= 0.5\n\n            epochs_without_improvement = 0\n\n    current_lrs = [\n        group[\"lr\"]\n        for group in optimizer.param_groups\n    ]\n\n    history.append(\n        {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"mean_f1\": mean_f1,\n            \"conv1_lr\": current_lrs[0],\n            \"layer4_lr\": current_lrs[1],\n            \"classifier_lr\": current_lrs[2]\n        }\n    )\n\n    print(\n        f\"\\nEpoch {epoch:02d}/{max_epochs}\"\n    )\n\n    print(\n        f\"Train Loss: {train_loss:.5f}\"\n    )\n\n    print(\n        f\"Val Loss:   {val_loss:.5f}\"\n    )\n\n    print(\n        f\"Mean F1:    {mean_f1:.4f}\"\n    )\n\n    print(\n        \"Learning Rates: \"\n        f\"Conv1={current_lrs[0]:.6f} | \"\n        f\"Layer4={current_lrs[1]:.6f} | \"\n        f\"Classifier={current_lrs[2]:.6f}\"\n    )\n\n    print(\n        f\"Best Epoch: {best_epoch}\"\n    )\n\n    print(\n        f\"Best Val Loss: {best_val_loss:.5f}\"\n    )\n\n    # ------------------------------------------------------------\n    # EARLY STOPPING\n    # ------------------------------------------------------------\n\n    if (\n        epoch - best_epoch\n        >= patience\n    ):\n        print(\n            \"\\nEarly stopping triggered.\"\n        )\n        break\n\n# ----------------------------------------------------------------\n# RESTORE BEST EXPERIMENTAL CHECKPOINT\n# ----------------------------------------------------------------\n\nif not os.path.exists(\n    experimental_checkpoint\n):\n    raise RuntimeError(\n        \"Experimental Fold-1 checkpoint was not saved.\"\n    )\n\nbest_state = torch.load(\n    experimental_checkpoint,\n    map_location=device\n)\n\nmodel.load_state_dict(\n    best_state,\n    strict=True\n)\n\nmodel.eval()\n\n# ----------------------------------------------------------------\n# SAVE HISTORY\n# ----------------------------------------------------------------\n\nhistory_df = pd.DataFrame(\n    history\n)\n\nhistory_df.to_csv(\n    history_path,\n    index=False\n)\n\nmetrics_df = pd.DataFrame(\n    {\n        \"fold\": [1],\n        \"best_epoch\": [best_epoch],\n        \"best_validation_loss\": [\n            best_val_loss\n        ]\n    }\n)\n\nmetrics_df.to_csv(\n    metrics_path,\n    index=False\n)\n\n# ----------------------------------------------------------------\n# FINAL VERIFICATION\n# ----------------------------------------------------------------\n\nif not torch.isfinite(\n    torch.cat(\n        [\n            parameter.detach().flatten()\n            for parameter in model.parameters()\n        ]\n    )\n).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in final model parameters.\"\n    )\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 54R - EXPERIMENT COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Best epoch: {best_epoch}\"\n)\n\nprint(\n    f\"Best validation loss: \"\n    f\"{best_val_loss:.6f}\"\n)\n\nprint(\n    f\"Experimental checkpoint saved: \"\n    f\"{experimental_checkpoint}\"\n)\n\nprint(\n    f\"History saved: {history_path}\"\n)\n\nprint(\n    f\"Metrics saved: {metrics_path}\"\n)\n\nprint(\n    \"\\nIMPORTANT:\"\n)\n\nprint(\n    \"Original Cell-40 Fold-1 checkpoint was NOT overwritten.\"\n)\n\nprint(\n    \"This is a controlled fine-tuning experiment only.\"\n)\n\nprint(\n    \"\\nCELL 54R COMPLETE\"\n)\n\nprint(\n    \"Send the complete Cell 54R output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:48:53.779405Z","iopub.execute_input":"2026-08-11T10:48:53.77977Z","iopub.status.idle":"2026-08-11T10:52:25.939632Z","shell.execute_reply.started":"2026-08-11T10:48:53.779741Z","shell.execute_reply":"2026-08-11T10:52:25.938458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 55 - FOLD 1 FINE-TUNED DIAGNOSTIC\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score\n\nprint(\"=\" * 70)\nprint(\"CELL 55 - FOLD 1 FINE-TUNED DIAGNOSTIC\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"val_loader_ft\",\n    \"TARGETS\",\n    \"device\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# CHECKPOINT\n# ----------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nfinetuned_checkpoint = os.path.join(\n    audit_dir,\n    \"cell54_fold1_finetuned_experiment.pt\"\n)\n\noutput_separation = os.path.join(\n    audit_dir,\n    \"cell55_fold1_finetuned_probability_separation.csv\"\n)\n\noutput_threshold = os.path.join(\n    audit_dir,\n    \"cell55_fold1_finetuned_threshold_audit.csv\"\n)\n\nif not os.path.exists(\n    finetuned_checkpoint\n):\n    raise RuntimeError(\n        \"Fine-tuned Fold-1 checkpoint not found: \"\n        + finetuned_checkpoint\n    )\n\nprint(\n    \"Fine-tuned Fold-1 checkpoint: EXISTS\"\n)\n\n# ----------------------------------------------------------------\n# LOAD BEST FINE-TUNED CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint = torch.load(\n    finetuned_checkpoint,\n    map_location=device\n)\n\nif isinstance(checkpoint, dict):\n\n    if \"state_dict\" in checkpoint:\n        checkpoint_state = checkpoint[\"state_dict\"]\n\n    elif \"model_state_dict\" in checkpoint:\n        checkpoint_state = checkpoint[\"model_state_dict\"]\n\n    else:\n        checkpoint_state = checkpoint\n\nelse:\n    raise RuntimeError(\n        \"Unsupported fine-tuned checkpoint format.\"\n    )\n\nmodel.load_state_dict(\n    checkpoint_state,\n    strict=True\n)\n\nmodel = model.to(device)\nmodel.eval()\n\nprint(\n    \"Fine-tuned Fold-1 checkpoint: LOADED\"\n)\n\nprint(\n    \"Model mode: evaluation\"\n)\n\n# ----------------------------------------------------------------\n# COLLECT VALIDATION PREDICTIONS\n# ----------------------------------------------------------------\n\nall_probabilities = []\nall_labels = []\nall_study_ids = []\n\nwith torch.no_grad():\n\n    for batch in val_loader_ft:\n\n        if not isinstance(batch, dict):\n            raise RuntimeError(\n                \"Unexpected validation batch type.\"\n            )\n\n        required_keys = [\n            \"image\",\n            \"label\",\n            \"study_id\"\n        ]\n\n        missing_keys = [\n            key\n            for key in required_keys\n            if key not in batch\n        ]\n\n        if missing_keys:\n            raise RuntimeError(\n                \"Missing validation batch keys: \"\n                + \", \".join(missing_keys)\n            )\n\n        images = batch[\"image\"].to(\n            device\n        )\n\n        labels = batch[\"label\"].to(\n            device\n        )\n\n        logits = model(\n            images\n        )\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n        if tuple(probabilities.shape) != (\n            images.shape[0],\n            len(TARGETS)\n        ):\n            raise RuntimeError(\n                \"Probability shape mismatch: \"\n                + str(tuple(probabilities.shape))\n            )\n\n        if not torch.isfinite(\n            probabilities\n        ).all():\n            raise RuntimeError(\n                \"NaN/Inf detected in probabilities.\"\n            )\n\n        all_probabilities.append(\n            probabilities.cpu()\n        )\n\n        all_labels.append(\n            labels.cpu()\n        )\n\n        all_study_ids.extend(\n            batch[\"study_id\"]\n        )\n\nprobabilities = torch.cat(\n    all_probabilities,\n    dim=0\n).numpy()\n\nlabels = torch.cat(\n    all_labels,\n    dim=0\n).numpy()\n\nprint(\n    f\"\\nProbability shape: {probabilities.shape}\"\n)\n\nprint(\n    f\"Label shape: {labels.shape}\"\n)\n\nif probabilities.shape != (\n    len(all_study_ids),\n    len(TARGETS)\n):\n    raise RuntimeError(\n        \"Final prediction shape is inconsistent.\"\n    )\n\nif labels.shape != probabilities.shape:\n    raise RuntimeError(\n        \"Probability/label shape mismatch.\"\n    )\n\nif len(set(all_study_ids)) != len(\n    all_study_ids\n):\n    raise RuntimeError(\n        \"Duplicate validation study predictions detected.\"\n    )\n\nprint(\n    \"Probability validity: PASS\"\n)\n\nprint(\n    \"Shape consistency: PASS\"\n)\n\nprint(\n    \"One prediction per validation study: PASS\"\n)\n\n# ----------------------------------------------------------------\n# TARGET-LEVEL DIAGNOSTICS\n# ----------------------------------------------------------------\n\nrows = []\n\nfor index, target in enumerate(TARGETS):\n\n    y_true = labels[:, index]\n    y_prob = probabilities[:, index]\n\n    positive_mask = (\n        y_true == 1\n    )\n\n    negative_mask = (\n        y_true == 0\n    )\n\n    positive_count = int(\n        positive_mask.sum()\n    )\n\n    negative_count = int(\n        negative_mask.sum()\n    )\n\n    if (\n        positive_count == 0\n        or negative_count == 0\n    ):\n        raise RuntimeError(\n            f\"Both classes are required for {target}.\"\n        )\n\n    positive_values = (\n        y_prob[positive_mask]\n    )\n\n    negative_values = (\n        y_prob[negative_mask]\n    )\n\n    positive_mean = float(\n        np.mean(positive_values)\n    )\n\n    negative_mean = float(\n        np.mean(negative_values)\n    )\n\n    positive_median = float(\n        np.median(positive_values)\n    )\n\n    negative_median = float(\n        np.median(negative_values)\n    )\n\n    probability_separation = (\n        positive_mean\n        - negative_mean\n    )\n\n    roc_auc = float(\n        roc_auc_score(\n            y_true,\n            y_prob\n        )\n    )\n\n    # ------------------------------------------------------------\n    # F1 @ 0.50\n    # ------------------------------------------------------------\n\n    predictions_050 = (\n        y_prob >= 0.50\n    ).astype(np.int32)\n\n    tp = np.sum(\n        (y_true == 1)\n        & (predictions_050 == 1)\n    )\n\n    fp = np.sum(\n        (y_true == 0)\n        & (predictions_050 == 1)\n    )\n\n    fn = np.sum(\n        (y_true == 1)\n        & (predictions_050 == 0)\n    )\n\n    denominator = (\n        2 * tp + fp + fn\n    )\n\n    if denominator == 0:\n        f1_at_050 = 0.0\n    else:\n        f1_at_050 = (\n            2 * tp\n            / denominator\n        )\n\n    # ------------------------------------------------------------\n    # BEST THRESHOLD\n    # ------------------------------------------------------------\n\n    thresholds = np.arange(\n        0.10,\n        0.9001,\n        0.025\n    )\n\n    best_threshold = 0.50\n    best_f1 = 0.0\n\n    for threshold in thresholds:\n\n        predictions = (\n            y_prob >= threshold\n        ).astype(np.int32)\n\n        tp_t = np.sum(\n            (y_true == 1)\n            & (predictions == 1)\n        )\n\n        fp_t = np.sum(\n            (y_true == 0)\n            & (predictions == 1)\n        )\n\n        fn_t = np.sum(\n            (y_true == 1)\n            & (predictions == 0)\n        )\n\n        denominator_t = (\n            2 * tp_t\n            + fp_t\n            + fn_t\n        )\n\n        if denominator_t == 0:\n            f1_t = 0.0\n        else:\n            f1_t = (\n                2 * tp_t\n                / denominator_t\n            )\n\n        if f1_t > best_f1:\n            best_f1 = float(f1_t)\n            best_threshold = float(\n                threshold\n            )\n\n    f1_improvement = (\n        best_f1\n        - f1_at_050\n    )\n\n    rows.append(\n        {\n            \"target\": target,\n            \"positive_count\": positive_count,\n            \"negative_count\": negative_count,\n            \"positive_mean\": positive_mean,\n            \"negative_mean\": negative_mean,\n            \"positive_median\": positive_median,\n            \"negative_median\": negative_median,\n            \"probability_separation\": probability_separation,\n            \"roc_auc\": roc_auc,\n            \"f1_at_0.50\": f1_at_050,\n            \"best_threshold\": best_threshold,\n            \"best_f1\": best_f1,\n            \"f1_improvement\": f1_improvement\n        }\n    )\n\ndiagnostics_df = pd.DataFrame(\n    rows\n)\n\n# ----------------------------------------------------------------\n# DISPLAY DIAGNOSTICS\n# ----------------------------------------------------------------\n\nprint(\"\\n\")\nprint(\n    diagnostics_df.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------\n\nmean_roc_auc = float(\n    diagnostics_df[\"roc_auc\"].mean()\n)\n\nmean_separation = float(\n    diagnostics_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    diagnostics_df[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nmean_best_f1 = float(\n    diagnostics_df[\n        \"best_f1\"\n    ].mean()\n)\n\npotential_f1_improvement = (\n    mean_best_f1\n    - mean_f1_050\n)\n\nroc_auc_good = diagnostics_df[\n    diagnostics_df[\"roc_auc\"] >= 0.65\n][\"target\"].tolist()\n\nnegative_separation = diagnostics_df[\n    diagnostics_df[\n        \"probability_separation\"\n    ] < 0\n][\"target\"].tolist()\n\nlarge_f1_improvement = diagnostics_df[\n    diagnostics_df[\n        \"f1_improvement\"\n    ] > 0.05\n][\"target\"].tolist()\n\nprint(\n    f\"\\nMean ROC-AUC: {mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50: \"\n    f\"{mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{potential_f1_improvement:.4f}\"\n)\n\nprint(\n    f\"\\nTargets with ROC-AUC >= 0.65: \"\n    f\"{len(roc_auc_good)}\"\n)\n\nprint(\n    roc_auc_good\n)\n\nprint(\n    f\"\\nTargets with negative probability separation: \"\n    f\"{len(negative_separation)}\"\n)\n\nprint(\n    negative_separation\n)\n\nprint(\n    f\"\\nTargets with >0.05 F1 improvement: \"\n    f\"{len(large_f1_improvement)}\"\n)\n\nprint(\n    large_f1_improvement\n)\n\n# ----------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------\n\ndiagnostics_df.to_csv(\n    output_separation,\n    index=False\n)\n\nthreshold_df = diagnostics_df[\n    [\n        \"target\",\n        \"f1_at_0.50\",\n        \"best_threshold\",\n        \"best_f1\",\n        \"f1_improvement\"\n    ]\n].copy()\n\nthreshold_df.to_csv(\n    output_threshold,\n    index=False\n)\n\nprint(\n    f\"\\nSaved: {output_separation}\"\n)\n\nprint(\n    f\"Saved: {output_threshold}\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No original checkpoint modified.\"\n)\n\nprint(\n    \"No test data used.\"\n)\n\nprint(\n    \"\\nCELL 55 COMPLETE\"\n)\n\nprint(\n    \"Send the complete Cell 55 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T10:54:23.707323Z","iopub.execute_input":"2026-08-11T10:54:23.707739Z","iopub.status.idle":"2026-08-11T10:54:35.833341Z","shell.execute_reply.started":"2026-08-11T10:54:23.707706Z","shell.execute_reply":"2026-08-11T10:54:35.832518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 56R - FOLD 2 FINE-TUNING FINALIZATION\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\n\nprint(\"=\" * 70)\nprint(\"CELL 56R - FOLD 2 FINE-TUNING FINALIZATION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\",\n    \"device\",\n    \"best_state_dict\",\n    \"best_epoch\",\n    \"best_val_loss\",\n    \"history\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# VERIFY BEST TRAINING RESULT\n# ----------------------------------------------------------------\n\nif best_epoch != 2:\n    raise RuntimeError(\n        f\"Unexpected best epoch: {best_epoch}\"\n    )\n\nif not np.isfinite(\n    float(best_val_loss)\n):\n    raise RuntimeError(\n        f\"Invalid best validation loss: {best_val_loss}\"\n    )\n\nprint(\n    f\"Best epoch from completed training: {best_epoch}\"\n)\n\nprint(\n    f\"Best validation loss from completed training: \"\n    f\"{float(best_val_loss):.6f}\"\n)\n\nprint(\n    \"Best training result: PASS\"\n)\n\n# ----------------------------------------------------------------\n# RESTORE BEST MODEL STATE\n# ----------------------------------------------------------------\n\nmodel.load_state_dict(\n    best_state_dict,\n    strict=True\n)\n\nmodel = model.to(device)\nmodel.eval()\n\nprint(\n    \"Best Fold-2 fine-tuned model state restored: PASS\"\n)\n\n# ----------------------------------------------------------------\n# MODEL VALIDITY CHECK\n# ----------------------------------------------------------------\n\nfor parameter in model.parameters():\n\n    if not torch.isfinite(\n        parameter\n    ).all():\n\n        raise RuntimeError(\n            \"NaN/Inf detected in restored model parameters.\"\n        )\n\nprint(\n    \"Restored model parameter validity: PASS\"\n)\n\n# ----------------------------------------------------------------\n# FORWARD-PASS CHECK\n# ----------------------------------------------------------------\n\nif \"val_loader\" in globals():\n\n    validation_batch = next(\n        iter(val_loader)\n    )\n\n    if not isinstance(\n        validation_batch,\n        dict\n    ):\n        raise RuntimeError(\n            \"Unexpected validation batch type.\"\n        )\n\n    required_keys = [\n        \"image\",\n        \"label\",\n        \"study_id\"\n    ]\n\n    missing_keys = [\n        key\n        for key in required_keys\n        if key not in validation_batch\n    ]\n\n    if missing_keys:\n        raise RuntimeError(\n            \"Missing validation batch keys: \"\n            + \", \".join(missing_keys)\n        )\n\n    images = validation_batch[\n        \"image\"\n    ].to(device)\n\n    labels = validation_batch[\n        \"label\"\n    ].to(device)\n\n    with torch.no_grad():\n\n        logits = model(\n            images\n        )\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n    if tuple(images.shape[1:]) != (\n        21,\n        224,\n        224\n    ):\n        raise RuntimeError(\n            \"Unexpected input shape: \"\n            + str(tuple(images.shape))\n        )\n\n    if tuple(logits.shape[1:]) != (\n        len(TARGETS),\n    ):\n        raise RuntimeError(\n            \"Unexpected output shape: \"\n            + str(tuple(logits.shape))\n        )\n\n    if not torch.isfinite(\n        logits\n    ).all():\n        raise RuntimeError(\n            \"NaN/Inf detected in logits.\"\n        )\n\n    if not torch.isfinite(\n        probabilities\n    ).all():\n        raise RuntimeError(\n            \"NaN/Inf detected in probabilities.\"\n        )\n\n    print(\n        \"Forward-pass validation: PASS\"\n    )\n\n# ----------------------------------------------------------------\n# SAVE EXPERIMENTAL CHECKPOINT\n# ----------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nfinetuned_checkpoint = os.path.join(\n    audit_dir,\n    \"cell56_fold2_finetuned_experiment.pt\"\n)\n\nhistory_path = os.path.join(\n    audit_dir,\n    \"cell56_fold2_finetuned_history.csv\"\n)\n\nmetrics_path = os.path.join(\n    audit_dir,\n    \"cell56_fold2_finetuned_metrics.csv\"\n)\n\ntorch.save(\n    {\n        \"state_dict\": model.state_dict(),\n        \"fold\": 2,\n        \"best_epoch\": int(best_epoch),\n        \"best_val_loss\": float(best_val_loss),\n        \"trainable_parameters\": 8_593_996,\n        \"config\": {\n            \"conv1_lr\": 1e-4,\n            \"layer4_lr\": 1e-5,\n            \"classifier_lr\": 1e-3,\n            \"weight_decay\": 1e-4,\n            \"gradient_clip\": 1.0,\n            \"max_epochs\": 8,\n            \"early_stopping_patience\": 3\n        }\n    },\n    finetuned_checkpoint\n)\n\nprint(\n    \"Experimental checkpoint saved: PASS\"\n)\n\n# ----------------------------------------------------------------\n# SAVE HISTORY\n# ----------------------------------------------------------------\n\nhistory_df = pd.DataFrame(\n    history\n)\n\nhistory_df.to_csv(\n    history_path,\n    index=False\n)\n\nprint(\n    \"History saved: PASS\"\n)\n\n# ----------------------------------------------------------------\n# SAVE METADATA\n# ----------------------------------------------------------------\n\nmetrics_df = pd.DataFrame(\n    [\n        {\n            \"fold\": 2,\n            \"best_epoch\": int(best_epoch),\n            \"best_val_loss\":\n                float(best_val_loss),\n            \"trainable_parameters\":\n                8_593_996,\n            \"conv1_lr\": 1e-4,\n            \"layer4_lr\": 1e-5,\n            \"classifier_lr\": 1e-3,\n            \"weight_decay\": 1e-4,\n            \"gradient_clip\": 1.0,\n            \"max_epochs\": 8,\n            \"early_stopping_patience\": 3\n        }\n    ]\n)\n\nmetrics_df.to_csv(\n    metrics_path,\n    index=False\n)\n\nprint(\n    \"Metrics saved: PASS\"\n)\n\n# ----------------------------------------------------------------\n# VERIFY FILES\n# ----------------------------------------------------------------\n\nfor path in [\n    finetuned_checkpoint,\n    history_path,\n    metrics_path\n]:\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            \"Expected output file was not created: \"\n            + path\n        )\n\nprint(\n    \"Output-file verification: PASS\"\n)\n\n# ----------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------\n\nprint(\"\\n\" + \"=\" * 70)\nprint(\"CELL 56R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Fold: 2\"\n)\n\nprint(\n    f\"Best epoch: {best_epoch}\"\n)\n\nprint(\n    f\"Best validation loss: \"\n    f\"{float(best_val_loss):.6f}\"\n)\n\nprint(\n    f\"Checkpoint: {finetuned_checkpoint}\"\n)\n\nprint(\n    \"Original Cell-44 Fold-2 checkpoint was NOT overwritten.\"\n)\n\nprint(\n    \"No additional training performed.\"\n)\n\nprint(\n    \"No test data used.\"\n)\n\nprint(\n    \"Fold-2 controlled fine-tuning experiment finalized successfully.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:01:17.217455Z","iopub.execute_input":"2026-08-11T11:01:17.217799Z","iopub.status.idle":"2026-08-11T11:01:18.356082Z","shell.execute_reply.started":"2026-08-11T11:01:17.217771Z","shell.execute_reply":"2026-08-11T11:01:18.355062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 57 - FOLD 2 FINE-TUNED DIAGNOSTIC\n# ================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.metrics import roc_auc_score\n\nprint(\"=\" * 70)\nprint(\"CELL 57 - FOLD 2 FINE-TUNED DIAGNOSTIC\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"val_loader\",\n    \"TARGETS\",\n    \"device\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell56_fold2_finetuned_experiment.pt\"\n)\n\nif not os.path.exists(\n    checkpoint_path\n):\n    raise RuntimeError(\n        \"Fine-tuned Fold-2 checkpoint not found: \"\n        + checkpoint_path\n    )\n\nprint(\n    \"Fine-tuned Fold-2 checkpoint: EXISTS\"\n)\n\n# ----------------------------------------------------------------\n# LOAD CHECKPOINT\n# ----------------------------------------------------------------\n\ncheckpoint = torch.load(\n    checkpoint_path,\n    map_location=device\n)\n\nif not isinstance(\n    checkpoint,\n    dict\n):\n    raise RuntimeError(\n        \"Unexpected checkpoint format.\"\n    )\n\nif \"state_dict\" not in checkpoint:\n    raise RuntimeError(\n        \"Checkpoint does not contain state_dict.\"\n    )\n\nmodel.load_state_dict(\n    checkpoint[\"state_dict\"],\n    strict=True\n)\n\nmodel = model.to(device)\nmodel.eval()\n\nprint(\n    \"Fine-tuned Fold-2 checkpoint: LOADED\"\n)\n\nprint(\n    \"Model mode: evaluation\"\n)\n\n# ----------------------------------------------------------------\n# COLLECT VALIDATION PREDICTIONS\n# ----------------------------------------------------------------\n\nprobability_batches = []\nlabel_batches = []\nstudy_ids = []\n\nwith torch.no_grad():\n\n    for batch in val_loader:\n\n        if not isinstance(\n            batch,\n            dict\n        ):\n            raise RuntimeError(\n                \"Unexpected validation batch type.\"\n            )\n\n        required_keys = [\n            \"image\",\n            \"label\",\n            \"study_id\"\n        ]\n\n        missing_keys = [\n            key\n            for key in required_keys\n            if key not in batch\n        ]\n\n        if missing_keys:\n            raise RuntimeError(\n                \"Missing validation batch keys: \"\n                + \", \".join(missing_keys)\n            )\n\n        images = batch[\n            \"image\"\n        ].to(device)\n\n        labels = batch[\n            \"label\"\n        ]\n\n        logits = model(\n            images\n        )\n\n        probabilities = torch.sigmoid(\n            logits\n        )\n\n        if not torch.isfinite(\n            probabilities\n        ).all():\n            raise RuntimeError(\n                \"NaN/Inf detected in probabilities.\"\n            )\n\n        probability_batches.append(\n            probabilities.cpu().numpy()\n        )\n\n        label_batches.append(\n            labels.cpu().numpy()\n        )\n\n        study_ids.extend(\n            list(batch[\"study_id\"])\n        )\n\nprobabilities = np.concatenate(\n    probability_batches,\n    axis=0\n)\n\nlabels = np.concatenate(\n    label_batches,\n    axis=0\n)\n\n# ----------------------------------------------------------------\n# SHAPE VALIDATION\n# ----------------------------------------------------------------\n\nif probabilities.shape != (\n    len(val_loader.dataset),\n    len(TARGETS)\n):\n    raise RuntimeError(\n        \"Probability shape mismatch: \"\n        + str(probabilities.shape)\n    )\n\nif labels.shape != (\n    len(val_loader.dataset),\n    len(TARGETS)\n):\n    raise RuntimeError(\n        \"Label shape mismatch: \"\n        + str(labels.shape)\n    )\n\nif len(study_ids) != len(\n    val_loader.dataset\n):\n    raise RuntimeError(\n        \"Study ID count does not match validation dataset.\"\n    )\n\nif len(set(study_ids)) != len(\n    study_ids\n):\n    raise RuntimeError(\n        \"Duplicate study IDs detected.\"\n    )\n\nprint(\n    f\"Probability shape: {probabilities.shape}\"\n)\n\nprint(\n    f\"Label shape: {labels.shape}\"\n)\n\nprint(\n    \"Probability validity: PASS\"\n)\n\nprint(\n    \"Shape consistency: PASS\"\n)\n\nprint(\n    \"One prediction per validation study: PASS\"\n)\n\n# ----------------------------------------------------------------\n# TARGET DIAGNOSTICS\n# ----------------------------------------------------------------\n\ndiagnostics = []\nthreshold_diagnostics = []\n\nthreshold_grid = np.arange(\n    0.10,\n    0.9001,\n    0.025\n)\n\nfor index, target in enumerate(\n    TARGETS\n):\n\n    y_true = labels[\n        :,\n        index\n    ].astype(int)\n\n    y_prob = probabilities[\n        :,\n        index\n    ]\n\n    positive_mask = (\n        y_true == 1\n    )\n\n    negative_mask = (\n        y_true == 0\n    )\n\n    positive_count = int(\n        positive_mask.sum()\n    )\n\n    negative_count = int(\n        negative_mask.sum()\n    )\n\n    positive_mean = float(\n        y_prob[\n            positive_mask\n        ].mean()\n    )\n\n    negative_mean = float(\n        y_prob[\n            negative_mask\n        ].mean()\n    )\n\n    positive_median = float(\n        np.median(\n            y_prob[\n                positive_mask\n            ]\n        )\n    )\n\n    negative_median = float(\n        np.median(\n            y_prob[\n                negative_mask\n            ]\n        )\n    )\n\n    probability_separation = (\n        positive_mean\n        - negative_mean\n    )\n\n    if (\n        positive_count > 0\n        and negative_count > 0\n    ):\n\n        roc_auc = float(\n            roc_auc_score(\n                y_true,\n                y_prob\n            )\n        )\n\n    else:\n\n        roc_auc = np.nan\n\n    # ------------------------------------------------------------\n    # F1 AT 0.50\n    # ------------------------------------------------------------\n\n    y_pred_050 = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    tp = int(\n        np.sum(\n            (y_true == 1)\n            & (y_pred_050 == 1)\n        )\n    )\n\n    fp = int(\n        np.sum(\n            (y_true == 0)\n            & (y_pred_050 == 1)\n        )\n    )\n\n    fn = int(\n        np.sum(\n            (y_true == 1)\n            & (y_pred_050 == 0)\n        )\n    )\n\n    f1_denominator = (\n        2 * tp\n        + fp\n        + fn\n    )\n\n    if f1_denominator == 0:\n        f1_at_050 = 0.0\n    else:\n        f1_at_050 = (\n            2 * tp\n            / f1_denominator\n        )\n\n    # ------------------------------------------------------------\n    # BEST THRESHOLD\n    # ------------------------------------------------------------\n\n    best_threshold = 0.50\n    best_f1 = f1_at_050\n\n    for threshold in threshold_grid:\n\n        y_pred = (\n            y_prob >= threshold\n        ).astype(int)\n\n        tp_t = int(\n            np.sum(\n                (y_true == 1)\n                & (y_pred == 1)\n            )\n        )\n\n        fp_t = int(\n            np.sum(\n                (y_true == 0)\n                & (y_pred == 1)\n            )\n        )\n\n        fn_t = int(\n            np.sum(\n                (y_true == 1)\n                & (y_pred == 0)\n            )\n        )\n\n        denominator = (\n            2 * tp_t\n            + fp_t\n            + fn_t\n        )\n\n        if denominator == 0:\n            f1_t = 0.0\n        else:\n            f1_t = (\n                2 * tp_t\n                / denominator\n            )\n\n        if f1_t > best_f1:\n            best_f1 = f1_t\n            best_threshold = float(\n                threshold\n            )\n\n    f1_improvement = (\n        best_f1\n        - f1_at_050\n    )\n\n    diagnostics.append(\n        {\n            \"target\": target,\n            \"positive_count\":\n                positive_count,\n            \"negative_count\":\n                negative_count,\n            \"positive_mean\":\n                positive_mean,\n            \"negative_mean\":\n                negative_mean,\n            \"positive_median\":\n                positive_median,\n            \"negative_median\":\n                negative_median,\n            \"probability_separation\":\n                probability_separation,\n            \"roc_auc\":\n                roc_auc,\n            \"f1_at_0.50\":\n                f1_at_050,\n            \"best_threshold\":\n                best_threshold,\n            \"best_f1\":\n                best_f1,\n            \"f1_improvement\":\n                f1_improvement\n        }\n    )\n\n    threshold_diagnostics.append(\n        {\n            \"target\": target,\n            \"f1_at_0.50\":\n                f1_at_050,\n            \"best_threshold\":\n                best_threshold,\n            \"best_f1\":\n                best_f1,\n            \"f1_improvement\":\n                f1_improvement\n        }\n    )\n\ndiagnostics_df = pd.DataFrame(\n    diagnostics\n)\n\nthreshold_df = pd.DataFrame(\n    threshold_diagnostics\n)\n\n# ----------------------------------------------------------------\n# DISPLAY\n# ----------------------------------------------------------------\n\nprint()\n\nprint(\n    diagnostics_df.to_string(\n        index=False,\n        float_format=lambda x:\n            f\"{x:.4f}\"\n    )\n)\n\nprint()\n\nprint(\n    threshold_df.to_string(\n        index=False,\n        float_format=lambda x:\n            f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------\n\nmean_roc_auc = float(\n    diagnostics_df[\n        \"roc_auc\"\n    ].mean()\n)\n\nmean_separation = float(\n    diagnostics_df[\n        \"probability_separation\"\n    ].mean()\n)\n\nmean_f1_050 = float(\n    diagnostics_df[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nmean_best_f1 = float(\n    diagnostics_df[\n        \"best_f1\"\n    ].mean()\n)\n\npotential_improvement = (\n    mean_best_f1\n    - mean_f1_050\n)\n\nhigh_auc_targets = diagnostics_df.loc[\n    diagnostics_df[\"roc_auc\"] >= 0.65,\n    \"target\"\n].tolist()\n\nnegative_separation_targets = diagnostics_df.loc[\n    diagnostics_df[\n        \"probability_separation\"\n    ] < 0,\n    \"target\"\n].tolist()\n\nimproved_targets = diagnostics_df.loc[\n    diagnostics_df[\n        \"f1_improvement\"\n    ] > 0.05,\n    \"target\"\n].tolist()\n\nprint()\n\nprint(\n    f\"Mean ROC-AUC: {mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"Mean probability separation: \"\n    f\"{mean_separation:.4f}\"\n)\n\nprint(\n    f\"Mean F1 @ 0.50: {mean_f1_050:.4f}\"\n)\n\nprint(\n    f\"Mean best-threshold F1: \"\n    f\"{mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Potential F1 improvement: \"\n    f\"{potential_improvement:.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Targets with ROC-AUC >= 0.65: \"\n    f\"{len(high_auc_targets)}\"\n)\n\nprint(\n    high_auc_targets\n)\n\nprint()\n\nprint(\n    f\"Targets with negative probability separation: \"\n    f\"{len(negative_separation_targets)}\"\n)\n\nprint(\n    negative_separation_targets\n)\n\nprint()\n\nprint(\n    f\"Targets with >0.05 F1 improvement: \"\n    f\"{len(improved_targets)}\"\n)\n\nprint(\n    improved_targets\n)\n\n# ----------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\ndiagnostics_path = os.path.join(\n    audit_dir,\n    \"cell57_fold2_finetuned_probability_separation.csv\"\n)\n\nthreshold_path = os.path.join(\n    audit_dir,\n    \"cell57_fold2_finetuned_threshold_audit.csv\"\n)\n\ndiagnostics_df.to_csv(\n    diagnostics_path,\n    index=False\n)\n\nthreshold_df.to_csv(\n    threshold_path,\n    index=False\n)\n\nprint()\n\nprint(\n    f\"Saved: {diagnostics_path}\"\n)\n\nprint(\n    f\"Saved: {threshold_path}\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No original checkpoint modified.\"\n)\n\nprint(\n    \"No test data used.\"\n)\n\nprint()\n\nprint(\"=\" * 70)\nprint(\"CELL 57 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Send the complete Cell 57 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:02:14.728607Z","iopub.execute_input":"2026-08-11T11:02:14.729Z","iopub.status.idle":"2026-08-11T11:02:23.934557Z","shell.execute_reply.started":"2026-08-11T11:02:14.728968Z","shell.execute_reply":"2026-08-11T11:02:23.933615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================================================\n# CELL 58A - FINE-TUNED CHECKPOINT FORMAT AUDIT\n# ================================================================\n\nimport os\nimport torch\n\nprint(\"=\" * 70)\nprint(\"CELL 58A - FINE-TUNED CHECKPOINT FORMAT AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\",\n    \"device\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------\n# CHECKPOINT PATHS\n# ----------------------------------------------------------------\n\ncheckpoint_paths = {\n    0: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell52_fold0_finetuned_experiment.pt\",\n\n    1: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell54_fold1_finetuned_experiment.pt\",\n\n    2: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell56_fold2_finetuned_experiment.pt\"\n}\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT FILE VALIDATION\")\nprint(\"=\" * 70)\n\nfor fold, path in checkpoint_paths.items():\n\n    exists = os.path.exists(path)\n\n    print(\n        f\"Fold-{fold} checkpoint: \"\n        f\"{'EXISTS' if exists else 'MISSING'}\"\n    )\n\n    if not exists:\n        raise RuntimeError(\n            f\"Missing Fold-{fold} fine-tuned checkpoint: \"\n            + path\n        )\n\nprint(\"All three fine-tuned checkpoints: PASS\")\n\n# ----------------------------------------------------------------\n# MODEL ARCHITECTURE VALIDATION\n# ----------------------------------------------------------------\n\nif not isinstance(\n    model,\n    torch.nn.Module\n):\n    raise RuntimeError(\n        \"Existing model is not a torch.nn.Module.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"EXISTING MODEL VALIDATION\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Model type: {type(model)}\"\n)\n\nprint(\n    f\"Model device: \"\n    f\"{next(model.parameters()).device}\"\n)\n\n# ----------------------------------------------------------------\n# CHECKPOINT FORMAT INSPECTION\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT FORMAT INSPECTION\")\nprint(\"=\" * 70)\n\ncheckpoint_audit = {}\n\nfor fold, path in checkpoint_paths.items():\n\n    print()\n    print(\n        f\"--- Fold {fold} ---\"\n    )\n\n    checkpoint = torch.load(\n        path,\n        map_location=\"cpu\",\n        weights_only=False\n    )\n\n    checkpoint_type = type(\n        checkpoint\n    )\n\n    print(\n        f\"Object type: {checkpoint_type}\"\n    )\n\n    if isinstance(\n        checkpoint,\n        dict\n    ):\n\n        print(\n            f\"Dictionary keys: \"\n            f\"{list(checkpoint.keys())}\"\n        )\n\n        for key, value in checkpoint.items():\n\n            print(\n                f\"  {key}: \"\n                f\"type={type(value)}\"\n            )\n\n            if isinstance(\n                value,\n                dict\n            ):\n                print(\n                    f\"       nested keys=\"\n                    f\"{list(value.keys())[:20]}\"\n                )\n\n    elif isinstance(\n        checkpoint,\n        torch.nn.Module\n    ):\n\n        print(\n            \"Checkpoint contains a complete \"\n            \"PyTorch model object.\"\n        )\n\n    else:\n\n        print(\n            \"Checkpoint is neither a standard \"\n            \"dictionary nor a torch.nn.Module.\"\n        )\n\n    checkpoint_audit[fold] = checkpoint\n\n# ----------------------------------------------------------------\n# DETERMINE LOADABLE STATE FORMAT WITHOUT ASSUMPTIONS\n# ----------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT LOAD FORMAT DETECTION\")\nprint(\"=\" * 70)\n\ndef detect_checkpoint_format(\n    checkpoint\n):\n\n    # Direct state_dict:\n    # dictionary whose values are tensors/parameters.\n    if isinstance(\n        checkpoint,\n        dict\n    ):\n\n        if len(checkpoint) > 0:\n\n            tensor_values = all(\n                isinstance(\n                    value,\n                    (\n                        torch.Tensor,\n                        torch.nn.Parameter\n                    )\n                )\n                for value in checkpoint.values()\n            )\n\n            if tensor_values:\n                return \"direct_state_dict\"\n\n        # Common wrapped formats.\n        for key in [\n            \"state_dict\",\n            \"model_state_dict\",\n            \"model\"\n        ]:\n\n            if key in checkpoint:\n\n                value = checkpoint[key]\n\n                if isinstance(\n                    value,\n                    dict\n                ):\n\n                    if len(value) == 0:\n                        continue\n\n                    tensor_values = all(\n                        isinstance(\n                            item,\n                            (\n                                torch.Tensor,\n                                torch.nn.Parameter\n                            )\n                        )\n                        for item in value.values()\n                    )\n\n                    if tensor_values:\n                        return (\n                            f\"wrapped_state_dict:{key}\"\n                        )\n\n                if isinstance(\n                    value,\n                    torch.nn.Module\n                ):\n                    return (\n                        f\"wrapped_model:{key}\"\n                    )\n\n    if isinstance(\n        checkpoint,\n        torch.nn.Module\n    ):\n        return \"complete_model\"\n\n    return \"unknown\"\n\ndetected_formats = {}\n\nfor fold, checkpoint in checkpoint_audit.items():\n\n    fmt = detect_checkpoint_format(\n        checkpoint\n    )\n\n    detected_formats[fold] = fmt\n\n    print(\n        f\"Fold-{fold}: {fmt}\"\n    )\n\n    if fmt == \"unknown\":\n\n        raise RuntimeError(\n            f\"Unable to safely determine \"\n            f\"Fold-{fold} checkpoint format. \"\n            f\"Do not continue.\"\n        )\n\n# ----------------------------------------------------------------\n# SAVE AUDIT IN NOTEBOOK STATE\n# ----------------------------------------------------------------\n\nfine_tuned_checkpoint_audit = {\n    \"paths\": checkpoint_paths,\n    \"formats\": detected_formats\n}\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 58A VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    \"All three checkpoint files: PASS\"\n)\n\nprint(\n    \"Checkpoint formats identified: PASS\"\n)\n\nfor fold in [0, 1, 2]:\n\n    print(\n        f\"Fold-{fold} format: \"\n        f\"{detected_formats[fold]}\"\n    )\n\nprint()\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No checkpoint modified.\"\n)\n\nprint(\n    \"No dataset modified.\"\n)\n\nprint(\n    \"No validation predictions generated.\"\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 58A COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Send the complete Cell 58A output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:08:54.287484Z","iopub.execute_input":"2026-08-11T11:08:54.287833Z","iopub.status.idle":"2026-08-11T11:08:54.392823Z","shell.execute_reply.started":"2026-08-11T11:08:54.2878Z","shell.execute_reply":"2026-08-11T11:08:54.391881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 58R - FINE-TUNED 3-FOLD OOF EVALUATION\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 70)\nprint(\"CELL 58R - FINE-TUNED 3-FOLD OOF EVALUATION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\",\n    \"KneeStudyDataset\",\n    \"modeling_table\",\n    \"study_folds\",\n    \"device\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Device: {device}\")\n\n# ----------------------------------------------------------------------\n# BASIC VALIDATION\n# ----------------------------------------------------------------------\n\nif not isinstance(model, torch.nn.Module):\n    raise RuntimeError(\n        \"Existing model is not a torch.nn.Module.\"\n    )\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nrequired_primary_columns = [\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\"\n]\n\nmissing_primary_columns = [\n    col\n    for col in required_primary_columns\n    if col not in modeling_table.columns\n]\n\nif missing_primary_columns:\n    raise RuntimeError(\n        \"Missing primary-series columns: \"\n        + \", \".join(missing_primary_columns)\n    )\n\nif \"StudyInstanceUID\" not in modeling_table.columns:\n    raise RuntimeError(\n        \"Missing StudyInstanceUID in modeling_table.\"\n    )\n\nif \"fold\" not in study_folds.columns:\n    raise RuntimeError(\n        \"Missing fold column in study_folds.\"\n    )\n\nprint(f\"Modeling table shape: {modeling_table.shape}\")\nprint(f\"Fold table shape: {study_folds.shape}\")\nprint(\"Target columns: PASS\")\nprint(\"Primary-series columns: PASS\")\n\n# ----------------------------------------------------------------------\n# BUILD / VERIFY FOLD ASSIGNMENT\n# ----------------------------------------------------------------------\n\nmodeling_ids = set(\n    modeling_table[\"StudyInstanceUID\"].astype(str)\n)\n\nfold_ids = set(\n    study_folds[\"StudyInstanceUID\"].astype(str)\n)\n\nif modeling_ids != fold_ids:\n    missing_fold_ids = modeling_ids - fold_ids\n    extra_fold_ids = fold_ids - modeling_ids\n\n    raise RuntimeError(\n        \"Modeling/fold study IDs do not match. \"\n        f\"Missing in fold table={len(missing_fold_ids)}, \"\n        f\"extra in fold table={len(extra_fold_ids)}.\"\n    )\n\nfold_assignment = study_folds[\n    [\"StudyInstanceUID\", \"fold\"]\n].copy()\n\nfold_assignment[\"StudyInstanceUID\"] = (\n    fold_assignment[\"StudyInstanceUID\"].astype(str)\n)\n\nfold_assignment[\"fold\"] = (\n    fold_assignment[\"fold\"].astype(int)\n)\n\nif fold_assignment[\"StudyInstanceUID\"].duplicated().any():\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID values in fold table.\"\n    )\n\noof_table_base = modeling_table.copy()\n\noof_table_base[\"StudyInstanceUID\"] = (\n    oof_table_base[\"StudyInstanceUID\"].astype(str)\n)\n\nif \"fold\" in oof_table_base.columns:\n    oof_table_base = oof_table_base.drop(\n        columns=[\"fold\"]\n    )\n\noof_table_base = oof_table_base.merge(\n    fold_assignment,\n    on=\"StudyInstanceUID\",\n    how=\"left\",\n    validate=\"one_to_one\"\n)\n\nif oof_table_base[\"fold\"].isna().any():\n    raise RuntimeError(\n        \"Some modeling studies have no fold assignment.\"\n    )\n\noof_table_base[\"fold\"] = (\n    oof_table_base[\"fold\"].astype(int)\n)\n\nprint(\"Fold assignment merge: PASS\")\nprint(\n    f\"Total modeling studies: \"\n    f\"{len(oof_table_base)}\"\n)\n\nprint(\n    f\"Unique studies: \"\n    f\"{oof_table_base['StudyInstanceUID'].nunique()}\"\n)\n\navailable_folds = sorted(\n    oof_table_base[\"fold\"].unique().tolist()\n)\n\nprint(\n    f\"Available folds: {available_folds}\"\n)\n\nif available_folds != [0, 1, 2]:\n    raise RuntimeError(\n        \"Expected exactly folds [0, 1, 2]. \"\n        f\"Found {available_folds}.\"\n    )\n\nprint(\"Three-fold configuration: PASS\")\n\n# ----------------------------------------------------------------------\n# FINE-TUNED CHECKPOINTS\n# ----------------------------------------------------------------------\n\ncheckpoint_paths = {\n    0: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell52_fold0_finetuned_experiment.pt\",\n\n    1: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell54_fold1_finetuned_experiment.pt\",\n\n    2: \"/kaggle/working/rsna_knee_audit/\"\n       \"cell56_fold2_finetuned_experiment.pt\"\n}\n\nfor fold, path in checkpoint_paths.items():\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Missing fine-tuned Fold-{fold} checkpoint: \"\n            + path\n        )\n\n    print(\n        f\"Fine-tuned Fold-{fold} checkpoint: EXISTS\"\n    )\n\nprint(\"All three fine-tuned checkpoints: PASS\")\n\n# ----------------------------------------------------------------------\n# CHECKPOINT STATE-DICT EXTRACTION\n# ----------------------------------------------------------------------\n\ndef extract_state_dict(checkpoint, fold):\n\n    # Fold 0 and Fold 1 were verified as direct state_dict objects.\n    if isinstance(\n        checkpoint,\n        dict\n    ):\n\n        # Direct state_dict:\n        # every value is a tensor/parameter.\n        if len(checkpoint) > 0:\n\n            direct_state_dict = all(\n                isinstance(\n                    value,\n                    (\n                        torch.Tensor,\n                        torch.nn.Parameter\n                    )\n                )\n                for value in checkpoint.values()\n            )\n\n            if direct_state_dict:\n                return checkpoint\n\n        # Fold 2 was verified as a wrapped state_dict.\n        if \"state_dict\" in checkpoint:\n\n            state_dict = checkpoint[\"state_dict\"]\n\n            if not isinstance(\n                state_dict,\n                dict\n            ):\n                raise RuntimeError(\n                    f\"Fold-{fold} state_dict is not a dictionary.\"\n                )\n\n            return state_dict\n\n    raise RuntimeError(\n        f\"Unable to extract state_dict from \"\n        f\"Fold-{fold} checkpoint.\"\n    )\n\n# ----------------------------------------------------------------------\n# MODEL STATE VALIDATION\n# ----------------------------------------------------------------------\n\nreference_model_state = model.state_dict()\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINE-TUNED CHECKPOINT COMPATIBILITY\")\nprint(\"=\" * 70)\n\nfor fold, path in checkpoint_paths.items():\n\n    checkpoint = torch.load(\n        path,\n        map_location=\"cpu\",\n        weights_only=False\n    )\n\n    state_dict = extract_state_dict(\n        checkpoint,\n        fold\n    )\n\n    checkpoint_keys = set(\n        state_dict.keys()\n    )\n\n    model_keys = set(\n        reference_model_state.keys()\n    )\n\n    missing_keys = model_keys - checkpoint_keys\n    unexpected_keys = checkpoint_keys - model_keys\n\n    if missing_keys:\n        raise RuntimeError(\n            f\"Fold-{fold} checkpoint missing model keys: \"\n            + \", \".join(\n                sorted(missing_keys)[:10]\n            )\n        )\n\n    if unexpected_keys:\n        raise RuntimeError(\n            f\"Fold-{fold} checkpoint has unexpected \"\n            f\"model keys: \"\n            + \", \".join(\n                sorted(unexpected_keys)[:10]\n            )\n        )\n\n    shape_mismatches = []\n\n    for key in model_keys:\n\n        if (\n            tuple(\n                state_dict[key].shape\n            )\n            !=\n            tuple(\n                reference_model_state[key].shape\n            )\n        ):\n            shape_mismatches.append(key)\n\n    if shape_mismatches:\n        raise RuntimeError(\n            f\"Fold-{fold} checkpoint has parameter \"\n            f\"shape mismatches: \"\n            + \", \".join(\n                shape_mismatches[:10]\n            )\n        )\n\n    print(\n        f\"Fold-{fold} state_dict compatibility: PASS\"\n    )\n\n# ----------------------------------------------------------------------\n# OOF INFERENCE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINE-TUNED OOF INFERENCE\")\nprint(\"=\" * 70)\n\noof_prediction_rows = []\n\nfor fold in [0, 1, 2]:\n\n    print()\n    print(\n        f\"FINE-TUNED FOLD {fold} OOF INFERENCE\"\n    )\n    print(\"-\" * 70)\n\n    fold_validation_table = oof_table_base[\n        oof_table_base[\"fold\"] == fold\n    ].copy()\n\n    fold_training_table = oof_table_base[\n        oof_table_base[\"fold\"] != fold\n    ].copy()\n\n    training_ids = set(\n        fold_training_table[\n            \"StudyInstanceUID\"\n        ].astype(str)\n    )\n\n    validation_ids = set(\n        fold_validation_table[\n            \"StudyInstanceUID\"\n        ].astype(str)\n    )\n\n    overlap = training_ids & validation_ids\n\n    print(\n        f\"Training studies: \"\n        f\"{len(training_ids)}\"\n    )\n\n    print(\n        f\"Validation studies: \"\n        f\"{len(validation_ids)}\"\n    )\n\n    print(\n        f\"Train/validation overlap: \"\n        f\"{len(overlap)}\"\n    )\n\n    if overlap:\n        raise RuntimeError(\n            f\"Fold-{fold} has train/validation \"\n            f\"study overlap.\"\n        )\n\n    # --------------------------------------------------------------\n    # EXACT EXISTING DATASET PIPELINE\n    # --------------------------------------------------------------\n\n    validation_dataset = KneeStudyDataset(\n        fold_validation_table,\n        targets=TARGETS\n    )\n\n    validation_loader = DataLoader(\n        validation_dataset,\n        batch_size=2,\n        shuffle=False,\n        num_workers=0,\n        pin_memory=False\n    )\n\n    print(\n        f\"Validation dataset: \"\n        f\"{len(validation_dataset)}\"\n    )\n\n    print(\n        f\"Validation batches: \"\n        f\"{len(validation_loader)}\"\n    )\n\n    if len(validation_dataset) != len(\n        fold_validation_table\n    ):\n        raise RuntimeError(\n            f\"Fold-{fold} validation dataset size \"\n            f\"does not match validation table.\"\n        )\n\n    # --------------------------------------------------------------\n    # LOAD CORRECT FINE-TUNED CHECKPOINT\n    # --------------------------------------------------------------\n\n    checkpoint = torch.load(\n        checkpoint_paths[fold],\n        map_location=\"cpu\",\n        weights_only=False\n    )\n\n    state_dict = extract_state_dict(\n        checkpoint,\n        fold\n    )\n\n    model.load_state_dict(\n        state_dict,\n        strict=True\n    )\n\n    model = model.to(device)\n    model.eval()\n\n    print(\n        f\"Fine-tuned Fold-{fold} checkpoint: LOADED\"\n    )\n\n    # --------------------------------------------------------------\n    # INFERENCE\n    # --------------------------------------------------------------\n\n    fold_prediction_count = 0\n    fold_study_ids = []\n\n    with torch.no_grad():\n\n        for batch in validation_loader:\n\n            if not isinstance(\n                batch,\n                dict\n            ):\n                raise RuntimeError(\n                    f\"Fold-{fold} batch is not a dictionary.\"\n                )\n\n            required_batch_keys = [\n                \"image\",\n                \"label\",\n                \"study_id\"\n            ]\n\n            missing_batch_keys = [\n                key\n                for key in required_batch_keys\n                if key not in batch\n            ]\n\n            if missing_batch_keys:\n                raise RuntimeError(\n                    f\"Fold-{fold} batch missing keys: \"\n                    + \", \".join(\n                        missing_batch_keys\n                    )\n                )\n\n            images = batch[\"image\"].to(\n                device,\n                non_blocking=True\n            )\n\n            labels = batch[\"label\"]\n\n            study_ids = batch[\"study_id\"]\n\n            if images.ndim != 4:\n                raise RuntimeError(\n                    f\"Unexpected image shape in \"\n                    f\"Fold-{fold}: \"\n                    f\"{tuple(images.shape)}\"\n                )\n\n            if images.shape[1:] != (\n                21,\n                224,\n                224\n            ):\n                raise RuntimeError(\n                    f\"Unexpected image shape in \"\n                    f\"Fold-{fold}: \"\n                    f\"{tuple(images.shape)}\"\n                )\n\n            if labels.ndim != 2:\n                raise RuntimeError(\n                    f\"Unexpected label shape in \"\n                    f\"Fold-{fold}: \"\n                    f\"{tuple(labels.shape)}\"\n                )\n\n            if labels.shape[1] != len(TARGETS):\n                raise RuntimeError(\n                    f\"Unexpected label target count \"\n                    f\"in Fold-{fold}: \"\n                    f\"{labels.shape[1]}\"\n                )\n\n            logits = model(images)\n\n            if logits.ndim != 2:\n                raise RuntimeError(\n                    f\"Unexpected logit shape in \"\n                    f\"Fold-{fold}: \"\n                    f\"{tuple(logits.shape)}\"\n                )\n\n            if logits.shape[1] != len(TARGETS):\n                raise RuntimeError(\n                    f\"Expected {len(TARGETS)} outputs, \"\n                    f\"got {logits.shape[1]}.\"\n                )\n\n            probabilities = torch.sigmoid(\n                logits\n            ).detach().cpu().numpy()\n\n            labels_np = (\n                labels.detach()\n                .cpu()\n                .numpy()\n            )\n\n            probabilities = np.asarray(\n                probabilities,\n                dtype=np.float64\n            )\n\n            labels_np = np.asarray(\n                labels_np,\n                dtype=np.float64\n            )\n\n            if not np.isfinite(\n                probabilities\n            ).all():\n                raise RuntimeError(\n                    f\"NaN/Inf probabilities in \"\n                    f\"Fold-{fold}.\"\n                )\n\n            if not np.isfinite(\n                labels_np\n            ).all():\n                raise RuntimeError(\n                    f\"NaN/Inf labels in \"\n                    f\"Fold-{fold}.\"\n                )\n\n            if (\n                probabilities.shape\n                != labels_np.shape\n            ):\n                raise RuntimeError(\n                    f\"Probability/label shape mismatch \"\n                    f\"in Fold-{fold}: \"\n                    f\"{probabilities.shape} vs \"\n                    f\"{labels_np.shape}\"\n                )\n\n            if len(study_ids) != len(\n                probabilities\n            ):\n                raise RuntimeError(\n                    f\"Study ID count does not match \"\n                    f\"prediction count in Fold-{fold}.\"\n                )\n\n            for row_idx, study_id in enumerate(\n                study_ids\n            ):\n\n                study_id = str(\n                    study_id\n                )\n\n                row = {\n                    \"StudyInstanceUID\": study_id,\n                    \"fold\": fold\n                }\n\n                for target_idx, target in enumerate(\n                    TARGETS\n                ):\n\n                    row[\n                        f\"{target}_prob\"\n                    ] = float(\n                        probabilities[\n                            row_idx,\n                            target_idx\n                        ]\n                    )\n\n                    row[\n                        target\n                    ] = float(\n                        labels_np[\n                            row_idx,\n                            target_idx\n                        ]\n                    )\n\n                oof_prediction_rows.append(\n                    row\n                )\n\n                fold_study_ids.append(\n                    study_id\n                )\n\n            fold_prediction_count += len(\n                probabilities\n            )\n\n    # --------------------------------------------------------------\n    # FOLD COVERAGE VALIDATION\n    # --------------------------------------------------------------\n\n    expected_ids = set(\n        fold_validation_table[\n            \"StudyInstanceUID\"\n        ].astype(str)\n    )\n\n    predicted_ids = set(\n        fold_study_ids\n    )\n\n    if expected_ids != predicted_ids:\n\n        missing_predictions = (\n            expected_ids - predicted_ids\n        )\n\n        unexpected_predictions = (\n            predicted_ids - expected_ids\n        )\n\n        raise RuntimeError(\n            f\"Fold-{fold} OOF coverage mismatch. \"\n            f\"Missing={len(missing_predictions)}, \"\n            f\"Unexpected={len(unexpected_predictions)}.\"\n        )\n\n    if len(fold_study_ids) != len(\n        expected_ids\n    ):\n        raise RuntimeError(\n            f\"Fold-{fold} contains duplicate \"\n            f\"validation predictions.\"\n        )\n\n    print(\n        f\"Fold-{fold} predictions: \"\n        f\"{fold_prediction_count}\"\n    )\n\n    print(\n        f\"Fold-{fold} coverage: PASS\"\n    )\n\n# ----------------------------------------------------------------------\n# BUILD FINAL OOF TABLE\n# ----------------------------------------------------------------------\n\nfine_tuned_oof = pd.DataFrame(\n    oof_prediction_rows\n)\n\nif fine_tuned_oof.empty:\n    raise RuntimeError(\n        \"Fine-tuned OOF prediction table is empty.\"\n    )\n\nexpected_oof_rows = len(\n    oof_table_base\n)\n\nif len(fine_tuned_oof) != expected_oof_rows:\n    raise RuntimeError(\n        f\"Expected {expected_oof_rows} OOF rows, \"\n        f\"got {len(fine_tuned_oof)}.\"\n    )\n\nif (\n    fine_tuned_oof[\"StudyInstanceUID\"]\n    .nunique()\n    != expected_oof_rows\n):\n    raise RuntimeError(\n        \"Fine-tuned OOF contains duplicate \"\n        \"StudyInstanceUID values.\"\n    )\n\nif (\n    fine_tuned_oof[\"fold\"]\n    .value_counts()\n    .sort_index()\n    .to_dict()\n    != {0: 20, 1: 19, 2: 19}\n):\n    raise RuntimeError(\n        \"Unexpected fine-tuned OOF fold distribution.\"\n    )\n\n# ----------------------------------------------------------------------\n# VALIDATE PROBABILITIES\n# ----------------------------------------------------------------------\n\nprobability_columns = [\n    f\"{target}_prob\"\n    for target in TARGETS\n]\n\nprobability_matrix = (\n    fine_tuned_oof[\n        probability_columns\n    ]\n    .to_numpy(\n        dtype=np.float64\n    )\n)\n\nlabel_matrix = (\n    fine_tuned_oof[\n        TARGETS\n    ]\n    .to_numpy(\n        dtype=np.float64\n    )\n)\n\nif not np.isfinite(\n    probability_matrix\n).all():\n    raise RuntimeError(\n        \"Fine-tuned OOF contains NaN/Inf probabilities.\"\n    )\n\nif not np.isfinite(\n    label_matrix\n).all():\n    raise RuntimeError(\n        \"Fine-tuned OOF contains NaN/Inf labels.\"\n    )\n\nif (\n    probability_matrix.min() < 0.0\n    or probability_matrix.max() > 1.0\n):\n    raise RuntimeError(\n        \"Fine-tuned OOF probabilities are outside [0, 1].\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINE-TUNED OOF VALIDATION\")\nprint(\"=\" * 70)\n\nprint(\n    f\"OOF shape: {fine_tuned_oof.shape}\"\n)\n\nprint(\n    f\"OOF rows: {len(fine_tuned_oof)}\"\n)\n\nprint(\n    f\"Unique OOF studies: \"\n    f\"{fine_tuned_oof['StudyInstanceUID'].nunique()}\"\n)\n\nprint(\n    \"One prediction per study: PASS\"\n)\n\nprint()\nprint(\"OOF fold distribution:\")\n\nprint(\n    fine_tuned_oof[\n        \"fold\"\n    ].value_counts()\n    .sort_index()\n)\n\nprint(\n    \"Fold coverage: PASS\"\n)\n\nprint(\n    \"Probability validity: PASS\"\n)\n\nprint(\n    \"Label validity: PASS\"\n)\n\n# ----------------------------------------------------------------------\n# SAVE OOF PREDICTIONS\n# ----------------------------------------------------------------------\n\noutput_path = (\n    \"/kaggle/working/rsna_knee_audit/\"\n    \"cell58_finetuned_oof_predictions.csv\"\n)\n\nfine_tuned_oof.to_csv(\n    output_path,\n    index=False\n)\n\nif not os.path.exists(\n    output_path\n):\n    raise RuntimeError(\n        \"Fine-tuned OOF output file was not created.\"\n    )\n\nprint()\nprint(\n    f\"Saved: {output_path}\"\n)\n\n# ----------------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 58R VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    \"Fine-tuned 3-fold OOF inference: PASS\"\n)\n\nprint(\n    \"58/58 studies covered exactly once: PASS\"\n)\n\nprint(\n    \"Fold distribution 20/19/19: PASS\"\n)\n\nprint(\n    \"21-channel input: PASS\"\n)\n\nprint(\n    \"12-target output: PASS\"\n)\n\nprint(\n    \"Checkpoint compatibility: PASS\"\n)\n\nprint(\n    \"Probability validity: PASS\"\n)\n\nprint(\n    \"No training performed.\"\n)\n\nprint(\n    \"No original checkpoint modified.\"\n)\n\nprint(\n    \"No test data used.\"\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 58R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Fine-tuned OOF predictions are ready \"\n    \"for comparison against baseline Cell-46 OOF.\"\n)\n\nprint(\n    \"Send the complete Cell 58R output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:14:12.918128Z","iopub.execute_input":"2026-08-11T11:14:12.918493Z","iopub.status.idle":"2026-08-11T11:14:45.324307Z","shell.execute_reply.started":"2026-08-11T11:14:12.918465Z","shell.execute_reply":"2026-08-11T11:14:45.323075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 59R - BASELINE vs FINE-TUNED OOF COMPARISON\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    f1_score,\n    brier_score_loss\n)\n\nprint(\"=\" * 70)\nprint(\"CELL 59R - BASELINE vs FINE-TUNED OOF COMPARISON\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECT\n# ----------------------------------------------------------------------\n\nif \"TARGETS\" not in globals():\n    raise RuntimeError(\n        \"Missing notebook object: TARGETS. Do not continue.\"\n    )\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Targets: {len(TARGETS)}\")\n\n# ----------------------------------------------------------------------\n# FILE PATHS\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nbaseline_path = os.path.join(\n    audit_dir,\n    \"cell46_oof_predictions.csv\"\n)\n\nfinetuned_path = os.path.join(\n    audit_dir,\n    \"cell58_finetuned_oof_predictions.csv\"\n)\n\nif not os.path.exists(baseline_path):\n    raise RuntimeError(\n        \"Baseline OOF file missing: \"\n        + baseline_path\n    )\n\nif not os.path.exists(finetuned_path):\n    raise RuntimeError(\n        \"Fine-tuned OOF file missing: \"\n        + finetuned_path\n    )\n\nprint(\"Baseline OOF file: EXISTS\")\nprint(\"Fine-tuned OOF file: EXISTS\")\n\n# ----------------------------------------------------------------------\n# LOAD FILES\n# ----------------------------------------------------------------------\n\nbaseline_oof = pd.read_csv(\n    baseline_path\n)\n\nfinetuned_oof = pd.read_csv(\n    finetuned_path\n)\n\nprint()\nprint(\n    f\"Baseline OOF shape: {baseline_oof.shape}\"\n)\n\nprint(\n    f\"Fine-tuned OOF shape: {finetuned_oof.shape}\"\n)\n\n# ----------------------------------------------------------------------\n# PRINT ACTUAL SCHEMA FIRST\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"ACTUAL BASELINE OOF SCHEMA\")\nprint(\"=\" * 70)\n\nfor i, column in enumerate(\n    baseline_oof.columns\n):\n    print(\n        f\"{i:02d}: {column}\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"ACTUAL FINE-TUNED OOF SCHEMA\")\nprint(\"=\" * 70)\n\nfor i, column in enumerate(\n    finetuned_oof.columns\n):\n    print(\n        f\"{i:02d}: {column}\"\n    )\n\n# ----------------------------------------------------------------------\n# IDENTIFY ID / FOLD COLUMNS\n# ----------------------------------------------------------------------\n\ndef find_column(\n    columns,\n    candidates\n):\n\n    for candidate in candidates:\n        if candidate in columns:\n            return candidate\n\n    return None\n\n\nbaseline_id_col = find_column(\n    baseline_oof.columns,\n    [\n        \"StudyInstanceUID\",\n        \"study_id\",\n        \"StudyID\"\n    ]\n)\n\nfinetuned_id_col = find_column(\n    finetuned_oof.columns,\n    [\n        \"StudyInstanceUID\",\n        \"study_id\",\n        \"StudyID\"\n    ]\n)\n\nif baseline_id_col is None:\n    raise RuntimeError(\n        \"Could not identify baseline study-ID column.\"\n    )\n\nif finetuned_id_col is None:\n    raise RuntimeError(\n        \"Could not identify fine-tuned study-ID column.\"\n    )\n\nif \"fold\" not in baseline_oof.columns:\n    raise RuntimeError(\n        \"Baseline OOF is missing the fold column.\"\n    )\n\nif \"fold\" not in finetuned_oof.columns:\n    raise RuntimeError(\n        \"Fine-tuned OOF is missing the fold column.\"\n    )\n\nprint()\nprint(\n    f\"Baseline study-ID column: {baseline_id_col}\"\n)\n\nprint(\n    f\"Fine-tuned study-ID column: {finetuned_id_col}\"\n)\n\nprint(\"Fold column: PASS\")\n\n# ----------------------------------------------------------------------\n# IDENTIFY PROBABILITY COLUMNS\n# ----------------------------------------------------------------------\n\ndef identify_probability_columns(\n    table\n):\n\n    columns = list(table.columns)\n\n    candidates = [\n        column\n        for column in columns\n        if column.endswith(\"_prob\")\n    ]\n\n    if len(candidates) != len(TARGETS):\n        raise RuntimeError(\n            \"Expected exactly \"\n            f\"{len(TARGETS)} probability columns, \"\n            f\"found {len(candidates)}: \"\n            + \", \".join(candidates)\n        )\n\n    return candidates\n\n\nbaseline_prob_columns = (\n    identify_probability_columns(\n        baseline_oof\n    )\n)\n\nfinetuned_prob_columns = (\n    identify_probability_columns(\n        finetuned_oof\n    )\n)\n\nprint()\nprint(\n    \"Baseline probability columns: \"\n    f\"{len(baseline_prob_columns)}\"\n)\n\nprint(\n    \"Fine-tuned probability columns: \"\n    f\"{len(finetuned_prob_columns)}\"\n)\n\nprint(\"Probability-column count: PASS\")\n\n# ----------------------------------------------------------------------\n# IDENTIFY LABEL COLUMNS FROM ACTUAL FILE STRUCTURE\n# ----------------------------------------------------------------------\n#\n# We do NOT assume that the label columns are named exactly\n# \"ACL\", \"MCL\", etc.\n#\n# The known structure is:\n#\n#   study ID\n#   fold\n#   12 label columns\n#   12 probability columns\n#\n# Therefore the 12 columns which are neither ID/fold/probability\n# columns are the label columns.\n# ----------------------------------------------------------------------\n\ndef identify_label_columns(\n    table,\n    id_col,\n    probability_columns\n):\n\n    excluded = {\n        id_col,\n        \"fold\"\n    }\n\n    excluded.update(\n        probability_columns\n    )\n\n    label_columns = [\n        column\n        for column in table.columns\n        if column not in excluded\n    ]\n\n    if len(label_columns) != len(TARGETS):\n        raise RuntimeError(\n            \"Could not identify exactly 12 label columns. \"\n            f\"Found {len(label_columns)}: \"\n            + \", \".join(label_columns)\n        )\n\n    return label_columns\n\n\nbaseline_label_columns = (\n    identify_label_columns(\n        baseline_oof,\n        baseline_id_col,\n        baseline_prob_columns\n    )\n)\n\nfinetuned_label_columns = (\n    identify_label_columns(\n        finetuned_oof,\n        finetuned_id_col,\n        finetuned_prob_columns\n    )\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"IDENTIFIED LABEL COLUMNS\")\nprint(\"=\" * 70)\n\nprint(\"Baseline labels:\")\n\nfor target, column in zip(\n    TARGETS,\n    baseline_label_columns\n):\n    print(\n        f\"{target:<22} -> {column}\"\n    )\n\nprint()\nprint(\"Fine-tuned labels:\")\n\nfor target, column in zip(\n    TARGETS,\n    finetuned_label_columns\n):\n    print(\n        f\"{target:<22} -> {column}\"\n    )\n\nprint()\nprint(\"Label-column count: PASS\")\n\n# ----------------------------------------------------------------------\n# BUILD CANONICAL TARGET MAPPINGS\n# ----------------------------------------------------------------------\n\nbaseline_target_map = {\n    target: column\n    for target, column in zip(\n        TARGETS,\n        baseline_label_columns\n    )\n}\n\nfinetuned_target_map = {\n    target: column\n    for target, column in zip(\n        TARGETS,\n        finetuned_label_columns\n    )\n}\n\nbaseline_probability_map = {\n    target: column\n    for target, column in zip(\n        TARGETS,\n        baseline_prob_columns\n    )\n}\n\nfinetuned_probability_map = {\n    target: column\n    for target, column in zip(\n        TARGETS,\n        finetuned_prob_columns\n    )\n}\n\nprint(\"Target-column mapping: PASS\")\n\n# ----------------------------------------------------------------------\n# STUDY ALIGNMENT\n# ----------------------------------------------------------------------\n\nbaseline_ids = (\n    baseline_oof[\n        baseline_id_col\n    ]\n    .astype(str)\n)\n\nfinetuned_ids = (\n    finetuned_oof[\n        finetuned_id_col\n    ]\n    .astype(str)\n)\n\nif len(baseline_oof) != 58:\n    raise RuntimeError(\n        \"Baseline OOF does not contain 58 rows.\"\n    )\n\nif len(finetuned_oof) != 58:\n    raise RuntimeError(\n        \"Fine-tuned OOF does not contain 58 rows.\"\n    )\n\nif baseline_ids.nunique() != 58:\n    raise RuntimeError(\n        \"Baseline OOF does not contain 58 unique studies.\"\n    )\n\nif finetuned_ids.nunique() != 58:\n    raise RuntimeError(\n        \"Fine-tuned OOF does not contain 58 unique studies.\"\n    )\n\nif set(baseline_ids) != set(finetuned_ids):\n    raise RuntimeError(\n        \"Baseline and fine-tuned OOF study sets differ.\"\n    )\n\n# ----------------------------------------------------------------------\n# CREATE SORTED COPIES FOR EXACT STUDY ALIGNMENT\n# ----------------------------------------------------------------------\n\nbaseline_eval = baseline_oof.copy()\n\nfinetuned_eval = finetuned_oof.copy()\n\nbaseline_eval[\"_study_key\"] = (\n    baseline_eval[\n        baseline_id_col\n    ].astype(str)\n)\n\nfinetuned_eval[\"_study_key\"] = (\n    finetuned_eval[\n        finetuned_id_col\n    ].astype(str)\n)\n\nbaseline_eval = (\n    baseline_eval\n    .sort_values(\"_study_key\")\n    .reset_index(drop=True)\n)\n\nfinetuned_eval = (\n    finetuned_eval\n    .sort_values(\"_study_key\")\n    .reset_index(drop=True)\n)\n\nif not (\n    baseline_eval[\"_study_key\"]\n    ==\n    finetuned_eval[\"_study_key\"]\n).all():\n    raise RuntimeError(\n        \"Study ordering/alignment failed.\"\n    )\n\n# ----------------------------------------------------------------------\n# FOLD ALIGNMENT\n# ----------------------------------------------------------------------\n\nif not (\n    baseline_eval[\"fold\"].astype(int)\n    ==\n    finetuned_eval[\"fold\"].astype(int)\n).all():\n    raise RuntimeError(\n        \"Baseline and fine-tuned fold assignments differ.\"\n    )\n\nfold_counts = (\n    baseline_eval[\n        \"fold\"\n    ]\n    .astype(int)\n    .value_counts()\n    .sort_index()\n    .to_dict()\n)\n\nexpected_fold_counts = {\n    0: 20,\n    1: 19,\n    2: 19\n}\n\nif fold_counts != expected_fold_counts:\n    raise RuntimeError(\n        \"Unexpected baseline fold distribution: \"\n        + str(fold_counts)\n    )\n\nfinetuned_fold_counts = (\n    finetuned_eval[\n        \"fold\"\n    ]\n    .astype(int)\n    .value_counts()\n    .sort_index()\n    .to_dict()\n)\n\nif finetuned_fold_counts != expected_fold_counts:\n    raise RuntimeError(\n        \"Unexpected fine-tuned fold distribution: \"\n        + str(finetuned_fold_counts)\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"OOF ALIGNMENT\")\nprint(\"=\" * 70)\n\nprint(\"Baseline studies: 58\")\nprint(\"Fine-tuned studies: 58\")\nprint(\"Same study set: PASS\")\nprint(\"Same fold assignments: PASS\")\nprint(\"Fold distribution: 20 / 19 / 19\")\nprint(\"OOF alignment: PASS\")\n\n# ----------------------------------------------------------------------\n# LABEL CONSISTENCY\n# ----------------------------------------------------------------------\n\nfor target in TARGETS:\n\n    baseline_labels = (\n        baseline_eval[\n            baseline_target_map[target]\n        ]\n        .to_numpy(\n            dtype=np.float64\n        )\n    )\n\n    finetuned_labels = (\n        finetuned_eval[\n            finetuned_target_map[target]\n        ]\n        .to_numpy(\n            dtype=np.float64\n        )\n    )\n\n    if not np.array_equal(\n        baseline_labels,\n        finetuned_labels\n    ):\n        raise RuntimeError(\n            f\"Label mismatch for target: {target}\"\n        )\n\nprint(\"Label consistency: PASS\")\n\n# ----------------------------------------------------------------------\n# PROBABILITY VALIDITY\n# ----------------------------------------------------------------------\n\nfor name, table, probability_map in [\n    (\n        \"Baseline\",\n        baseline_eval,\n        baseline_probability_map\n    ),\n    (\n        \"Fine-tuned\",\n        finetuned_eval,\n        finetuned_probability_map\n    )\n]:\n\n    probability_matrix = np.column_stack(\n        [\n            table[\n                probability_map[target]\n            ].to_numpy(\n                dtype=np.float64\n            )\n            for target in TARGETS\n        ]\n    )\n\n    if not np.isfinite(\n        probability_matrix\n    ).all():\n        raise RuntimeError(\n            f\"{name} probabilities contain NaN/Inf.\"\n        )\n\n    if (\n        probability_matrix.min() < 0.0\n        or probability_matrix.max() > 1.0\n    ):\n        raise RuntimeError(\n            f\"{name} probabilities outside [0,1].\"\n        )\n\nprint(\"Baseline probability validity: PASS\")\nprint(\"Fine-tuned probability validity: PASS\")\n\n# ----------------------------------------------------------------------\n# THRESHOLD GRID\n# ----------------------------------------------------------------------\n\nthresholds = np.linspace(\n    0.10,\n    0.90,\n    33\n)\n\n# ----------------------------------------------------------------------\n# METRIC FUNCTION\n# ----------------------------------------------------------------------\n\ndef calculate_metrics(\n    table,\n    target,\n    label_map,\n    probability_map\n):\n\n    y_true = (\n        table[\n            label_map[target]\n        ]\n        .to_numpy(\n            dtype=np.float64\n        )\n    )\n\n    y_prob = (\n        table[\n            probability_map[target]\n        ]\n        .to_numpy(\n            dtype=np.float64\n        )\n    )\n\n    positive_count = int(\n        y_true.sum()\n    )\n\n    negative_count = int(\n        len(y_true) - positive_count\n    )\n\n    if positive_count == 0:\n        raise RuntimeError(\n            f\"{target}: no positive OOF samples.\"\n        )\n\n    if negative_count == 0:\n        raise RuntimeError(\n            f\"{target}: no negative OOF samples.\"\n        )\n\n    roc_auc = roc_auc_score(\n        y_true,\n        y_prob\n    )\n\n    pr_auc = average_precision_score(\n        y_true,\n        y_prob\n    )\n\n    brier = brier_score_loss(\n        y_true,\n        y_prob\n    )\n\n    pred_050 = (\n        y_prob >= 0.50\n    ).astype(int)\n\n    f1_050 = f1_score(\n        y_true,\n        pred_050,\n        zero_division=0\n    )\n\n    best_threshold = None\n    best_f1 = -1.0\n\n    for threshold in thresholds:\n\n        predictions = (\n            y_prob >= threshold\n        ).astype(int)\n\n        current_f1 = f1_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n\n        if (\n            current_f1 > best_f1\n            or (\n                current_f1 == best_f1\n                and (\n                    best_threshold is None\n                    or threshold > best_threshold\n                )\n            )\n        ):\n            best_f1 = current_f1\n            best_threshold = float(\n                threshold\n            )\n\n    positive_mean = float(\n        y_prob[y_true == 1].mean()\n    )\n\n    negative_mean = float(\n        y_prob[y_true == 0].mean()\n    )\n\n    separation = (\n        positive_mean\n        -\n        negative_mean\n    )\n\n    return {\n        \"target\": target,\n        \"positive_count\": positive_count,\n        \"negative_count\": negative_count,\n        \"roc_auc\": float(roc_auc),\n        \"pr_auc\": float(pr_auc),\n        \"brier_score\": float(brier),\n        \"f1_at_0.50\": float(f1_050),\n        \"best_threshold\": float(\n            best_threshold\n        ),\n        \"best_f1\": float(best_f1),\n        \"f1_improvement\": float(\n            best_f1 - f1_050\n        ),\n        \"probability_separation\": float(\n            separation\n        )\n    }\n\n# ----------------------------------------------------------------------\n# CALCULATE METRICS\n# ----------------------------------------------------------------------\n\nbaseline_metrics = pd.DataFrame(\n    [\n        calculate_metrics(\n            baseline_eval,\n            target,\n            baseline_target_map,\n            baseline_probability_map\n        )\n        for target in TARGETS\n    ]\n)\n\nfinetuned_metrics = pd.DataFrame(\n    [\n        calculate_metrics(\n            finetuned_eval,\n            target,\n            finetuned_target_map,\n            finetuned_probability_map\n        )\n        for target in TARGETS\n    ]\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"BASELINE OOF METRICS\")\nprint(\"=\" * 70)\n\nprint(\n    baseline_metrics.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINE-TUNED OOF METRICS\")\nprint(\"=\" * 70)\n\nprint(\n    finetuned_metrics.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# COMPARISON TABLE\n# ----------------------------------------------------------------------\n\ncomparison = pd.DataFrame(\n    {\n        \"target\": TARGETS,\n\n        \"baseline_roc_auc\":\n            baseline_metrics[\n                \"roc_auc\"\n            ],\n\n        \"finetuned_roc_auc\":\n            finetuned_metrics[\n                \"roc_auc\"\n            ],\n\n        \"baseline_pr_auc\":\n            baseline_metrics[\n                \"pr_auc\"\n            ],\n\n        \"finetuned_pr_auc\":\n            finetuned_metrics[\n                \"pr_auc\"\n            ],\n\n        \"baseline_brier\":\n            baseline_metrics[\n                \"brier_score\"\n            ],\n\n        \"finetuned_brier\":\n            finetuned_metrics[\n                \"brier_score\"\n            ],\n\n        \"baseline_f1_at_0.50\":\n            baseline_metrics[\n                \"f1_at_0.50\"\n            ],\n\n        \"finetuned_f1_at_0.50\":\n            finetuned_metrics[\n                \"f1_at_0.50\"\n            ],\n\n        \"baseline_best_f1\":\n            baseline_metrics[\n                \"best_f1\"\n            ],\n\n        \"finetuned_best_f1\":\n            finetuned_metrics[\n                \"best_f1\"\n            ],\n\n        \"baseline_best_threshold\":\n            baseline_metrics[\n                \"best_threshold\"\n            ],\n\n        \"finetuned_best_threshold\":\n            finetuned_metrics[\n                \"best_threshold\"\n            ],\n\n        \"baseline_probability_separation\":\n            baseline_metrics[\n                \"probability_separation\"\n            ],\n\n        \"finetuned_probability_separation\":\n            finetuned_metrics[\n                \"probability_separation\"\n            ]\n    }\n)\n\ncomparison[\n    \"delta_roc_auc\"\n] = (\n    comparison[\n        \"finetuned_roc_auc\"\n    ]\n    -\n    comparison[\n        \"baseline_roc_auc\"\n    ]\n)\n\ncomparison[\n    \"delta_pr_auc\"\n] = (\n    comparison[\n        \"finetuned_pr_auc\"\n    ]\n    -\n    comparison[\n        \"baseline_pr_auc\"\n    ]\n)\n\ncomparison[\n    \"delta_brier\"\n] = (\n    comparison[\n        \"finetuned_brier\"\n    ]\n    -\n    comparison[\n        \"baseline_brier\"\n    ]\n)\n\ncomparison[\n    \"delta_f1_at_0.50\"\n] = (\n    comparison[\n        \"finetuned_f1_at_0.50\"\n    ]\n    -\n    comparison[\n        \"baseline_f1_at_0.50\"\n    ]\n)\n\ncomparison[\n    \"delta_best_f1\"\n] = (\n    comparison[\n        \"finetuned_best_f1\"\n    ]\n    -\n    comparison[\n        \"baseline_best_f1\"\n    ]\n)\n\ncomparison[\n    \"delta_probability_separation\"\n] = (\n    comparison[\n        \"finetuned_probability_separation\"\n    ]\n    -\n    comparison[\n        \"baseline_probability_separation\"\n    ]\n)\n\n# ----------------------------------------------------------------------\n# SUMMARY STATISTICS\n# ----------------------------------------------------------------------\n\nbaseline_mean_roc_auc = float(\n    baseline_metrics[\n        \"roc_auc\"\n    ].mean()\n)\n\nfinetuned_mean_roc_auc = float(\n    finetuned_metrics[\n        \"roc_auc\"\n    ].mean()\n)\n\nbaseline_mean_pr_auc = float(\n    baseline_metrics[\n        \"pr_auc\"\n    ].mean()\n)\n\nfinetuned_mean_pr_auc = float(\n    finetuned_metrics[\n        \"pr_auc\"\n    ].mean()\n)\n\nbaseline_mean_brier = float(\n    baseline_metrics[\n        \"brier_score\"\n    ].mean()\n)\n\nfinetuned_mean_brier = float(\n    finetuned_metrics[\n        \"brier_score\"\n    ].mean()\n)\n\nbaseline_mean_f1 = float(\n    baseline_metrics[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nfinetuned_mean_f1 = float(\n    finetuned_metrics[\n        \"f1_at_0.50\"\n    ].mean()\n)\n\nbaseline_mean_best_f1 = float(\n    baseline_metrics[\n        \"best_f1\"\n    ].mean()\n)\n\nfinetuned_mean_best_f1 = float(\n    finetuned_metrics[\n        \"best_f1\"\n    ].mean()\n)\n\nbaseline_mean_separation = float(\n    baseline_metrics[\n        \"probability_separation\"\n    ].mean()\n)\n\nfinetuned_mean_separation = float(\n    finetuned_metrics[\n        \"probability_separation\"\n    ].mean()\n)\n\n# ----------------------------------------------------------------------\n# PRINT COMPARISON\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"BASELINE vs FINE-TUNED OOF COMPARISON\")\nprint(\"=\" * 70)\n\nprint(\n    comparison[\n        [\n            \"target\",\n            \"baseline_roc_auc\",\n            \"finetuned_roc_auc\",\n            \"delta_roc_auc\",\n            \"baseline_f1_at_0.50\",\n            \"finetuned_f1_at_0.50\",\n            \"delta_f1_at_0.50\",\n            \"baseline_best_f1\",\n            \"finetuned_best_f1\",\n            \"delta_best_f1\"\n        ]\n    ].to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# MEAN SUMMARY\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"MEAN OOF SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Baseline mean ROC-AUC:       \"\n    f\"{baseline_mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean ROC-AUC:     \"\n    f\"{finetuned_mean_roc_auc:.4f}\"\n)\n\nprint(\n    f\"ROC-AUC delta:               \"\n    f\"{finetuned_mean_roc_auc - baseline_mean_roc_auc:+.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Baseline mean PR-AUC:        \"\n    f\"{baseline_mean_pr_auc:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean PR-AUC:      \"\n    f\"{finetuned_mean_pr_auc:.4f}\"\n)\n\nprint(\n    f\"PR-AUC delta:                \"\n    f\"{finetuned_mean_pr_auc - baseline_mean_pr_auc:+.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Baseline mean Brier:         \"\n    f\"{baseline_mean_brier:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean Brier:       \"\n    f\"{finetuned_mean_brier:.4f}\"\n)\n\nprint(\n    f\"Brier delta:                 \"\n    f\"{finetuned_mean_brier - baseline_mean_brier:+.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Baseline mean F1 @ 0.50:     \"\n    f\"{baseline_mean_f1:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean F1 @ 0.50:   \"\n    f\"{finetuned_mean_f1:.4f}\"\n)\n\nprint(\n    f\"F1 @ 0.50 delta:             \"\n    f\"{finetuned_mean_f1 - baseline_mean_f1:+.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Baseline mean best F1:       \"\n    f\"{baseline_mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean best F1:     \"\n    f\"{finetuned_mean_best_f1:.4f}\"\n)\n\nprint(\n    f\"Best F1 delta:               \"\n    f\"{finetuned_mean_best_f1 - baseline_mean_best_f1:+.4f}\"\n)\n\nprint()\n\nprint(\n    f\"Baseline mean probability separation: \"\n    f\"{baseline_mean_separation:.4f}\"\n)\n\nprint(\n    f\"Fine-tuned mean probability separation: \"\n    f\"{finetuned_mean_separation:.4f}\"\n)\n\nprint(\n    f\"Probability separation delta: \"\n    f\"{finetuned_mean_separation - baseline_mean_separation:+.4f}\"\n)\n\n# ----------------------------------------------------------------------\n# WIN COUNTS\n# ----------------------------------------------------------------------\n\nroc_auc_wins = int(\n    (\n        comparison[\n            \"delta_roc_auc\"\n        ] > 0\n    ).sum()\n)\n\npr_auc_wins = int(\n    (\n        comparison[\n            \"delta_pr_auc\"\n        ] > 0\n    ).sum()\n)\n\nf1_wins = int(\n    (\n        comparison[\n            \"delta_f1_at_0.50\"\n        ] > 0\n    ).sum()\n)\n\nbest_f1_wins = int(\n    (\n        comparison[\n            \"delta_best_f1\"\n        ] > 0\n    ).sum()\n)\n\nbrier_wins = int(\n    (\n        comparison[\n            \"delta_brier\"\n        ] < 0\n    ).sum()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"TARGET-LEVEL IMPROVEMENT COUNTS\")\nprint(\"=\" * 70)\n\nprint(\n    f\"ROC-AUC improvement:       \"\n    f\"{roc_auc_wins}/12\"\n)\n\nprint(\n    f\"PR-AUC improvement:        \"\n    f\"{pr_auc_wins}/12\"\n)\n\nprint(\n    f\"F1 @ 0.50 improvement:     \"\n    f\"{f1_wins}/12\"\n)\n\nprint(\n    f\"Best-F1 improvement:       \"\n    f\"{best_f1_wins}/12\"\n)\n\nprint(\n    f\"Brier improvement:         \"\n    f\"{brier_wins}/12\"\n)\n\n# ----------------------------------------------------------------------\n# SAVE RESULTS\n# ----------------------------------------------------------------------\n\ncomparison_path = os.path.join(\n    audit_dir,\n    \"cell59_baseline_vs_finetuned_oof.csv\"\n)\n\nbaseline_summary_path = os.path.join(\n    audit_dir,\n    \"cell59_baseline_oof_summary.csv\"\n)\n\nfinetuned_summary_path = os.path.join(\n    audit_dir,\n    \"cell59_finetuned_oof_summary.csv\"\n)\n\ncomparison.to_csv(\n    comparison_path,\n    index=False\n)\n\nbaseline_metrics.to_csv(\n    baseline_summary_path,\n    index=False\n)\n\nfinetuned_metrics.to_csv(\n    finetuned_summary_path,\n    index=False\n)\n\nif not all(\n    os.path.exists(path)\n    for path in [\n        comparison_path,\n        baseline_summary_path,\n        finetuned_summary_path\n    ]\n):\n    raise RuntimeError(\n        \"One or more Cell-59 output files were not created.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"OUTPUT FILES\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Saved: {comparison_path}\"\n)\n\nprint(\n    f\"Saved: {baseline_summary_path}\"\n)\n\nprint(\n    f\"Saved: {finetuned_summary_path}\"\n)\n\n# ----------------------------------------------------------------------\n# FINAL VERIFICATION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 59R VERIFICATION\")\nprint(\"=\" * 70)\n\nprint(\"58-study baseline OOF: PASS\")\nprint(\"58-study fine-tuned OOF: PASS\")\nprint(\"Same study set: PASS\")\nprint(\"Same fold assignments: PASS\")\nprint(\"12 target labels aligned: PASS\")\nprint(\"12 probability outputs aligned: PASS\")\nprint(\"Probability validity: PASS\")\nprint(\"Baseline metrics calculated: PASS\")\nprint(\"Fine-tuned metrics calculated: PASS\")\nprint(\"Comparison table: PASS\")\n\nprint()\nprint(\"No training performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No test data used.\")\nprint(\"No final model selected automatically.\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 59R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Send the complete Cell 59R output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:20:07.284167Z","iopub.execute_input":"2026-08-11T11:20:07.284512Z","iopub.status.idle":"2026-08-11T11:20:09.182046Z","shell.execute_reply.started":"2026-08-11T11:20:07.284485Z","shell.execute_reply":"2026-08-11T11:20:09.180833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 60 - OOF MODEL SELECTION + THRESHOLD DECISION AUDIT\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 60 - OOF MODEL SELECTION + THRESHOLD DECISION AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------------\n\nif \"TARGETS\" not in globals():\n    raise RuntimeError(\n        \"Missing notebook object: TARGETS. Do not continue.\"\n    )\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(f\"Targets: {len(TARGETS)}\")\n\n# ----------------------------------------------------------------------\n# PATHS\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\ncomparison_path = os.path.join(\n    audit_dir,\n    \"cell59_baseline_vs_finetuned_oof.csv\"\n)\n\nbaseline_summary_path = os.path.join(\n    audit_dir,\n    \"cell59_baseline_oof_summary.csv\"\n)\n\nfinetuned_summary_path = os.path.join(\n    audit_dir,\n    \"cell59_finetuned_oof_summary.csv\"\n)\n\nrequired_paths = [\n    comparison_path,\n    baseline_summary_path,\n    finetuned_summary_path\n]\n\nfor path in required_paths:\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Required Cell-59 output missing: {path}\"\n        )\n\nprint(\"Cell-59 comparison file: EXISTS\")\nprint(\"Cell-59 baseline summary: EXISTS\")\nprint(\"Cell-59 fine-tuned summary: EXISTS\")\n\n# ----------------------------------------------------------------------\n# LOAD RESULTS\n# ----------------------------------------------------------------------\n\ncomparison = pd.read_csv(\n    comparison_path\n)\n\nbaseline_metrics = pd.read_csv(\n    baseline_summary_path\n)\n\nfinetuned_metrics = pd.read_csv(\n    finetuned_summary_path\n)\n\n# ----------------------------------------------------------------------\n# SHAPE VALIDATION\n# ----------------------------------------------------------------------\n\nif comparison.shape[0] != 12:\n    raise RuntimeError(\n        f\"Expected 12 comparison rows, found {comparison.shape[0]}.\"\n    )\n\nif baseline_metrics.shape[0] != 12:\n    raise RuntimeError(\n        f\"Expected 12 baseline rows, found {baseline_metrics.shape[0]}.\"\n    )\n\nif finetuned_metrics.shape[0] != 12:\n    raise RuntimeError(\n        f\"Expected 12 fine-tuned rows, found {finetuned_metrics.shape[0]}.\"\n    )\n\nprint(\"Comparison rows: 12\")\nprint(\"Baseline metric rows: 12\")\nprint(\"Fine-tuned metric rows: 12\")\nprint(\"Target coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# TARGET ORDER VALIDATION\n# ----------------------------------------------------------------------\n\nif list(comparison[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Comparison target order does not match TARGETS.\"\n    )\n\nif list(baseline_metrics[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Baseline target order does not match TARGETS.\"\n    )\n\nif list(finetuned_metrics[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Fine-tuned target order does not match TARGETS.\"\n    )\n\nprint(\"Target ordering: PASS\")\n\n# ----------------------------------------------------------------------\n# REQUIRED COLUMNS\n# ----------------------------------------------------------------------\n\nrequired_comparison_columns = [\n    \"target\",\n    \"baseline_roc_auc\",\n    \"finetuned_roc_auc\",\n    \"delta_roc_auc\",\n    \"baseline_pr_auc\",\n    \"finetuned_pr_auc\",\n    \"delta_pr_auc\",\n    \"baseline_brier\",\n    \"finetuned_brier\",\n    \"delta_brier\",\n    \"baseline_f1_at_0.50\",\n    \"finetuned_f1_at_0.50\",\n    \"delta_f1_at_0.50\",\n    \"baseline_best_f1\",\n    \"finetuned_best_f1\",\n    \"delta_best_f1\",\n    \"baseline_probability_separation\",\n    \"finetuned_probability_separation\",\n    \"delta_probability_separation\"\n]\n\nmissing_columns = [\n    column\n    for column in required_comparison_columns\n    if column not in comparison.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Cell-59 comparison missing columns: \"\n        + \", \".join(missing_columns)\n    )\n\nprint(\"Comparison schema: PASS\")\n\n# ----------------------------------------------------------------------\n# CREATE MODEL DECISION\n# ----------------------------------------------------------------------\n#\n# We do not select a model from a single metric.\n#\n# Priority:\n#   1. ROC-AUC\n#   2. PR-AUC\n#   3. Best-F1\n#\n# Brier is treated separately because lower is better.\n#\n# A model is considered clearly better only when it improves at\n# least two of the three primary discrimination metrics.\n# ----------------------------------------------------------------------\n\ndecision_rows = []\n\nfor target in TARGETS:\n\n    row = comparison[\n        comparison[\"target\"] == target\n    ].iloc[0]\n\n    baseline_roc = float(\n        row[\"baseline_roc_auc\"]\n    )\n\n    finetuned_roc = float(\n        row[\"finetuned_roc_auc\"]\n    )\n\n    baseline_pr = float(\n        row[\"baseline_pr_auc\"]\n    )\n\n    finetuned_pr = float(\n        row[\"finetuned_pr_auc\"]\n    )\n\n    baseline_best_f1 = float(\n        row[\"baseline_best_f1\"]\n    )\n\n    finetuned_best_f1 = float(\n        row[\"finetuned_best_f1\"]\n    )\n\n    baseline_brier = float(\n        row[\"baseline_brier\"]\n    )\n\n    finetuned_brier = float(\n        row[\"finetuned_brier\"]\n    )\n\n    roc_improved = (\n        finetuned_roc > baseline_roc\n    )\n\n    pr_improved = (\n        finetuned_pr > baseline_pr\n    )\n\n    best_f1_improved = (\n        finetuned_best_f1 > baseline_best_f1\n    )\n\n    brier_improved = (\n        finetuned_brier < baseline_brier\n    )\n\n    primary_wins = int(roc_improved) + int(\n        pr_improved\n    ) + int(\n        best_f1_improved\n    )\n\n    primary_losses = int(\n        finetuned_roc < baseline_roc\n    ) + int(\n        finetuned_pr < baseline_pr\n    ) + int(\n        finetuned_best_f1 < baseline_best_f1\n    )\n\n    if primary_wins >= 2 and primary_losses == 0:\n        decision = \"fine_tuned\"\n\n    elif primary_losses >= 2 and primary_wins == 0:\n        decision = \"baseline\"\n\n    else:\n        decision = \"uncertain\"\n\n    decision_rows.append(\n        {\n            \"target\": target,\n            \"baseline_roc_auc\": baseline_roc,\n            \"finetuned_roc_auc\": finetuned_roc,\n            \"delta_roc_auc\": finetuned_roc - baseline_roc,\n            \"baseline_pr_auc\": baseline_pr,\n            \"finetuned_pr_auc\": finetuned_pr,\n            \"delta_pr_auc\": finetuned_pr - baseline_pr,\n            \"baseline_best_f1\": baseline_best_f1,\n            \"finetuned_best_f1\": finetuned_best_f1,\n            \"delta_best_f1\": (\n                finetuned_best_f1\n                -\n                baseline_best_f1\n            ),\n            \"baseline_brier\": baseline_brier,\n            \"finetuned_brier\": finetuned_brier,\n            \"delta_brier\": (\n                finetuned_brier\n                -\n                baseline_brier\n            ),\n            \"roc_auc_improved\": roc_improved,\n            \"pr_auc_improved\": pr_improved,\n            \"best_f1_improved\": best_f1_improved,\n            \"brier_improved\": brier_improved,\n            \"primary_metric_wins\": primary_wins,\n            \"primary_metric_losses\": primary_losses,\n            \"decision\": decision\n        }\n    )\n\ndecision_table = pd.DataFrame(\n    decision_rows\n)\n\n# ----------------------------------------------------------------------\n# THRESHOLD SELECTION\n# ----------------------------------------------------------------------\n\nthreshold_rows = []\n\nfor target in TARGETS:\n\n    baseline_row = baseline_metrics[\n        baseline_metrics[\"target\"] == target\n    ].iloc[0]\n\n    finetuned_row = finetuned_metrics[\n        finetuned_metrics[\"target\"] == target\n    ].iloc[0]\n\n    decision_row = decision_table[\n        decision_table[\"target\"] == target\n    ].iloc[0]\n\n    if decision_row[\"decision\"] == \"baseline\":\n\n        selected_model = \"baseline\"\n\n        selected_threshold = float(\n            baseline_row[\"best_threshold\"]\n        )\n\n        selected_best_f1 = float(\n            baseline_row[\"best_f1\"]\n        )\n\n    elif decision_row[\"decision\"] == \"fine_tuned\":\n\n        selected_model = \"fine_tuned\"\n\n        selected_threshold = float(\n            finetuned_row[\"best_threshold\"]\n        )\n\n        selected_best_f1 = float(\n            finetuned_row[\"best_f1\"]\n        )\n\n    else:\n\n        # For uncertain targets, choose the model with the\n        # higher best-F1 while keeping the decision explicitly\n        # marked as evidence-based rather than automatically\n        # declaring the architecture globally superior.\n\n        if (\n            float(finetuned_row[\"best_f1\"])\n            >\n            float(baseline_row[\"best_f1\"])\n        ):\n\n            selected_model = \"fine_tuned\"\n\n            selected_threshold = float(\n                finetuned_row[\"best_threshold\"]\n            )\n\n            selected_best_f1 = float(\n                finetuned_row[\"best_f1\"]\n            )\n\n        else:\n\n            selected_model = \"baseline\"\n\n            selected_threshold = float(\n                baseline_row[\"best_threshold\"]\n            )\n\n            selected_best_f1 = float(\n                baseline_row[\"best_f1\"]\n            )\n\n    threshold_rows.append(\n        {\n            \"target\": target,\n            \"decision\": decision_row[\"decision\"],\n            \"selected_model\": selected_model,\n            \"selected_threshold\": selected_threshold,\n            \"selected_best_f1\": selected_best_f1,\n            \"baseline_threshold\": float(\n                baseline_row[\"best_threshold\"]\n            ),\n            \"finetuned_threshold\": float(\n                finetuned_row[\"best_threshold\"]\n            )\n        }\n    )\n\nselection_table = pd.DataFrame(\n    threshold_rows\n)\n\n# ----------------------------------------------------------------------\n# PRINT MODEL DECISION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TARGET-LEVEL MODEL DECISION\")\nprint(\"=\" * 70)\n\nprint(\n    decision_table[\n        [\n            \"target\",\n            \"delta_roc_auc\",\n            \"delta_pr_auc\",\n            \"delta_best_f1\",\n            \"delta_brier\",\n            \"primary_metric_wins\",\n            \"primary_metric_losses\",\n            \"decision\"\n        ]\n    ].to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# PRINT THRESHOLD SELECTION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"OOF THRESHOLD SELECTION\")\nprint(\"=\" * 70)\n\nprint(\n    selection_table.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# SUMMARY\n# ----------------------------------------------------------------------\n\nbaseline_count = int(\n    (\n        decision_table[\"decision\"]\n        == \"baseline\"\n    ).sum()\n)\n\nfinetuned_count = int(\n    (\n        decision_table[\"decision\"]\n        == \"fine_tuned\"\n    ).sum()\n)\n\nuncertain_count = int(\n    (\n        decision_table[\"decision\"]\n        == \"uncertain\"\n    ).sum()\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"MODEL DECISION SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Baseline clearly preferred:     \"\n    f\"{baseline_count}/12\"\n)\n\nprint(\n    f\"Fine-tuned clearly preferred:   \"\n    f\"{finetuned_count}/12\"\n)\n\nprint(\n    f\"Uncertain targets:              \"\n    f\"{uncertain_count}/12\"\n)\n\n# ----------------------------------------------------------------------\n# SANITY CHECKS\n# ----------------------------------------------------------------------\n\nif not selection_table[\n    \"selected_threshold\"\n].between(\n    0.10,\n    0.90\n).all():\n\n    raise RuntimeError(\n        \"Selected threshold outside audited range.\"\n    )\n\nif not selection_table[\n    \"selected_model\"\n].isin(\n    [\n        \"baseline\",\n        \"fine_tuned\"\n    ]\n).all():\n\n    raise RuntimeError(\n        \"Invalid selected model.\"\n    )\n\nprint()\nprint(\"Threshold range validation: PASS\")\nprint(\"Model-selection validation: PASS\")\n\n# ----------------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------------\n\ndecision_path = os.path.join(\n    audit_dir,\n    \"cell60_model_selection_audit.csv\"\n)\n\nthreshold_path = os.path.join(\n    audit_dir,\n    \"cell60_selected_thresholds.csv\"\n)\n\ndecision_table.to_csv(\n    decision_path,\n    index=False\n)\n\nselection_table.to_csv(\n    threshold_path,\n    index=False\n)\n\nif not os.path.exists(\n    decision_path\n):\n    raise RuntimeError(\n        \"Model-selection audit was not saved.\"\n    )\n\nif not os.path.exists(\n    threshold_path\n):\n    raise RuntimeError(\n        \"Threshold selection was not saved.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"OUTPUT FILES\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Saved: {decision_path}\"\n)\n\nprint(\n    f\"Saved: {threshold_path}\"\n)\n\n# ----------------------------------------------------------------------\n# FINAL VERIFICATION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 60 VERIFICATION\")\nprint(\"=\" * 70)\n\nprint(\"12 targets evaluated: PASS\")\nprint(\"Baseline OOF evidence used: PASS\")\nprint(\"Fine-tuned OOF evidence used: PASS\")\nprint(\"Thresholds derived from OOF: PASS\")\nprint(\"No training performed: PASS\")\nprint(\"No checkpoint modified: PASS\")\nprint(\"No test data used: PASS\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 60 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Send the complete Cell 60 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:31:57.24894Z","iopub.execute_input":"2026-08-11T11:31:57.249316Z","iopub.status.idle":"2026-08-11T11:31:57.322457Z","shell.execute_reply.started":"2026-08-11T11:31:57.249284Z","shell.execute_reply":"2026-08-11T11:31:57.321612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 61 - FINAL OOF-BACKED MODEL + THRESHOLD PLAN AUDIT\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 61 - FINAL OOF-BACKED MODEL + THRESHOLD PLAN AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------------\n\nif \"TARGETS\" not in globals():\n    raise RuntimeError(\n        \"Missing notebook object: TARGETS. Do not continue.\"\n    )\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nprint(\"TARGETS: PASS\")\nprint(f\"Target count: {len(TARGETS)}\")\n\n# ----------------------------------------------------------------------\n# AUDIT DIRECTORY\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nif not os.path.isdir(audit_dir):\n    raise RuntimeError(\n        f\"Audit directory missing: {audit_dir}\"\n    )\n\nprint(\"Audit directory: PASS\")\n\n# ----------------------------------------------------------------------\n# REQUIRED CELL-60 OUTPUTS\n# ----------------------------------------------------------------------\n\nmodel_selection_path = os.path.join(\n    audit_dir,\n    \"cell60_model_selection_audit.csv\"\n)\n\nthreshold_selection_path = os.path.join(\n    audit_dir,\n    \"cell60_selected_thresholds.csv\"\n)\n\nrequired_files = [\n    model_selection_path,\n    threshold_selection_path\n]\n\nfor path in required_files:\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Required Cell-60 output missing: {path}\"\n        )\n\nprint(\"Cell-60 model-selection audit: EXISTS\")\nprint(\"Cell-60 threshold-selection audit: EXISTS\")\n\n# ----------------------------------------------------------------------\n# LOAD CELL-60 RESULTS\n# ----------------------------------------------------------------------\n\nmodel_selection = pd.read_csv(\n    model_selection_path\n)\n\nthreshold_selection = pd.read_csv(\n    threshold_selection_path\n)\n\n# ----------------------------------------------------------------------\n# BASIC VALIDATION\n# ----------------------------------------------------------------------\n\nif len(model_selection) != 12:\n    raise RuntimeError(\n        \"Cell-60 model-selection table must contain exactly 12 targets.\"\n    )\n\nif len(threshold_selection) != 12:\n    raise RuntimeError(\n        \"Cell-60 threshold table must contain exactly 12 targets.\"\n    )\n\nprint(\"Model-selection rows: 12\")\nprint(\"Threshold-selection rows: 12\")\n\n# ----------------------------------------------------------------------\n# TARGET ORDER\n# ----------------------------------------------------------------------\n\nif list(model_selection[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Model-selection target order does not match TARGETS.\"\n    )\n\nif list(threshold_selection[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Threshold-selection target order does not match TARGETS.\"\n    )\n\nprint(\"Target ordering: PASS\")\n\n# ----------------------------------------------------------------------\n# REQUIRED COLUMNS\n# ----------------------------------------------------------------------\n\nrequired_model_columns = [\n    \"target\",\n    \"decision\",\n    \"primary_metric_wins\",\n    \"primary_metric_losses\"\n]\n\nrequired_threshold_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"selected_best_f1\",\n    \"baseline_threshold\",\n    \"finetuned_threshold\"\n]\n\nmissing_model_columns = [\n    c for c in required_model_columns\n    if c not in model_selection.columns\n]\n\nmissing_threshold_columns = [\n    c for c in required_threshold_columns\n    if c not in threshold_selection.columns\n]\n\nif missing_model_columns:\n    raise RuntimeError(\n        \"Cell-60 model-selection table missing columns: \"\n        + \", \".join(missing_model_columns)\n    )\n\nif missing_threshold_columns:\n    raise RuntimeError(\n        \"Cell-60 threshold table missing columns: \"\n        + \", \".join(missing_threshold_columns)\n    )\n\nprint(\"Model-selection schema: PASS\")\nprint(\"Threshold-selection schema: PASS\")\n\n# ----------------------------------------------------------------------\n# VALID MODEL NAMES\n# ----------------------------------------------------------------------\n\nvalid_models = {\n    \"baseline\",\n    \"fine_tuned\"\n}\n\ninvalid_models = [\n    value\n    for value in threshold_selection[\"selected_model\"]\n    if value not in valid_models\n]\n\nif invalid_models:\n    raise RuntimeError(\n        \"Invalid selected model values: \"\n        + \", \".join(sorted(set(invalid_models)))\n    )\n\nprint(\"Selected model names: PASS\")\n\n# ----------------------------------------------------------------------\n# VALID THRESHOLDS\n# ----------------------------------------------------------------------\n\nthreshold_values = pd.to_numeric(\n    threshold_selection[\"selected_threshold\"],\n    errors=\"coerce\"\n)\n\nif threshold_values.isna().any():\n    raise RuntimeError(\n        \"Selected threshold contains NaN or non-numeric values.\"\n    )\n\nif not threshold_values.between(\n    0.10,\n    0.90\n).all():\n    raise RuntimeError(\n        \"Selected threshold outside audited range 0.10-0.90.\"\n    )\n\nprint(\"Selected threshold range: PASS\")\n\n# ----------------------------------------------------------------------\n# MODEL DECISION DISTRIBUTION\n# ----------------------------------------------------------------------\n\nbaseline_targets = threshold_selection.loc[\n    threshold_selection[\"selected_model\"] == \"baseline\",\n    \"target\"\n].tolist()\n\nfinetuned_targets = threshold_selection.loc[\n    threshold_selection[\"selected_model\"] == \"fine_tuned\",\n    \"target\"\n].tolist()\n\nif len(baseline_targets) + len(finetuned_targets) != 12:\n    raise RuntimeError(\n        \"Model-selection coverage does not equal 12 targets.\"\n    )\n\nif set(baseline_targets).intersection(\n    set(finetuned_targets)\n):\n    raise RuntimeError(\n        \"A target was assigned to both baseline and fine-tuned models.\"\n    )\n\nprint(\"Target model coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# IMPORTANT OOF DECISION CHECK\n# ----------------------------------------------------------------------\n#\n# We preserve the Cell-60 decision exactly.\n#\n# We do NOT overwrite uncertain decisions.\n# We do NOT invent new thresholds.\n# We do NOT retrain here.\n# ----------------------------------------------------------------------\n\nmerged_plan = threshold_selection[\n    [\n        \"target\",\n        \"selected_model\",\n        \"selected_threshold\",\n        \"selected_best_f1\"\n    ]\n].copy()\n\nmerged_plan[\"selected_threshold\"] = pd.to_numeric(\n    merged_plan[\"selected_threshold\"]\n)\n\nmerged_plan[\"selected_best_f1\"] = pd.to_numeric(\n    merged_plan[\"selected_best_f1\"]\n)\n\n# ----------------------------------------------------------------------\n# FINAL PLAN TABLE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL OOF-BACKED TARGET PLAN\")\nprint(\"=\" * 70)\n\nprint(\n    merged_plan.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# MODEL COUNTS\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL MODEL ASSIGNMENT SUMMARY\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Baseline-selected targets:   \"\n    f\"{len(baseline_targets)}/12\"\n)\n\nprint(\n    f\"Fine-tuned-selected targets: \"\n    f\"{len(finetuned_targets)}/12\"\n)\n\nprint()\nprint(\"Baseline targets:\")\nprint(baseline_targets)\n\nprint()\nprint(\"Fine-tuned targets:\")\nprint(finetuned_targets)\n\n# ----------------------------------------------------------------------\n# CHECKPOINT AVAILABILITY\n# ----------------------------------------------------------------------\n\nbaseline_checkpoint_paths = {\n    0: os.path.join(\n        audit_dir,\n        \"cell37_fold0_best_model.pt\"\n    ),\n    1: os.path.join(\n        audit_dir,\n        \"cell40_fold1_best_model.pt\"\n    ),\n    2: os.path.join(\n        audit_dir,\n        \"cell44_fold2_best_model.pt\"\n    )\n}\n\nfinetuned_checkpoint_paths = {\n    0: os.path.join(\n        audit_dir,\n        \"cell52_fold0_finetuned_experiment.pt\"\n    ),\n    1: os.path.join(\n        audit_dir,\n        \"cell54_fold1_finetuned_experiment.pt\"\n    ),\n    2: os.path.join(\n        audit_dir,\n        \"cell56_fold2_finetuned_experiment.pt\"\n    )\n}\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT AVAILABILITY\")\nprint(\"=\" * 70)\n\nfor fold, path in baseline_checkpoint_paths.items():\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Missing baseline Fold-{fold} checkpoint: {path}\"\n        )\n\n    print(\n        f\"Baseline Fold-{fold} checkpoint: EXISTS\"\n    )\n\nfor fold, path in finetuned_checkpoint_paths.items():\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Missing fine-tuned Fold-{fold} checkpoint: {path}\"\n        )\n\n    print(\n        f\"Fine-tuned Fold-{fold} checkpoint: EXISTS\"\n    )\n\nprint(\"All six validated checkpoints: PASS\")\n\n# ----------------------------------------------------------------------\n# SAVE FINAL PLAN\n# ----------------------------------------------------------------------\n\nfinal_plan_path = os.path.join(\n    audit_dir,\n    \"cell61_final_oof_backed_model_plan.csv\"\n)\n\nmerged_plan.to_csv(\n    final_plan_path,\n    index=False\n)\n\nif not os.path.exists(final_plan_path):\n    raise RuntimeError(\n        \"Final OOF-backed model plan was not saved.\"\n    )\n\n# ----------------------------------------------------------------------\n# FINAL VERIFICATION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 61 VERIFICATION\")\nprint(\"=\" * 70)\n\nprint(\"12 targets covered: PASS\")\nprint(\"OOF-derived model decisions preserved: PASS\")\nprint(\"OOF-derived thresholds preserved: PASS\")\nprint(\"Baseline checkpoints verified: PASS\")\nprint(\"Fine-tuned checkpoints verified: PASS\")\nprint(\"No new training performed: PASS\")\nprint(\"No checkpoint modified: PASS\")\nprint(\"No test data used: PASS\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"OUTPUT\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Saved: {final_plan_path}\"\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 61 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Final OOF-backed model plan is ready.\"\n)\n\nprint(\n    \"Send the complete Cell 61 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:33:06.217536Z","iopub.execute_input":"2026-08-11T11:33:06.217914Z","iopub.status.idle":"2026-08-11T11:33:06.263371Z","shell.execute_reply.started":"2026-08-11T11:33:06.217863Z","shell.execute_reply":"2026-08-11T11:33:06.262098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 62 - FINAL CHECKPOINT LOADING + INFERENCE ARCHITECTURE AUDIT\n# ======================================================================\n\nimport os\nimport copy\nimport torch\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 70)\nprint(\"CELL 62 - FINAL CHECKPOINT LOADING + INFERENCE ARCHITECTURE AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"model\",\n    \"TARGETS\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------------\n\nif \"device\" in globals():\n    inference_device = device\nelse:\n    inference_device = torch.device(\n        \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    )\n\nprint(f\"Device: {inference_device}\")\n\n# ----------------------------------------------------------------------\n# BASIC MODEL VALIDATION\n# ----------------------------------------------------------------------\n\nif not isinstance(model, torch.nn.Module):\n    raise RuntimeError(\n        \"Existing model object is not a torch.nn.Module.\"\n    )\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nmodel = model.to(inference_device)\n\nprint(\"Existing model: PASS\")\nprint(\"12-target configuration: PASS\")\n\n# ----------------------------------------------------------------------\n# AUDIT DIRECTORY\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nif not os.path.isdir(audit_dir):\n    raise RuntimeError(\n        f\"Audit directory missing: {audit_dir}\"\n    )\n\nprint(\"Audit directory: PASS\")\n\n# ----------------------------------------------------------------------\n# FINAL PLAN\n# ----------------------------------------------------------------------\n\nplan_path = os.path.join(\n    audit_dir,\n    \"cell61_final_oof_backed_model_plan.csv\"\n)\n\nif not os.path.exists(plan_path):\n    raise RuntimeError(\n        f\"Final model plan missing: {plan_path}\"\n    )\n\nfinal_plan = pd.read_csv(plan_path)\n\nif len(final_plan) != 12:\n    raise RuntimeError(\n        \"Final model plan must contain exactly 12 targets.\"\n    )\n\nrequired_plan_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"selected_best_f1\"\n]\n\nmissing_plan_columns = [\n    c\n    for c in required_plan_columns\n    if c not in final_plan.columns\n]\n\nif missing_plan_columns:\n    raise RuntimeError(\n        \"Final model plan missing columns: \"\n        + \", \".join(missing_plan_columns)\n    )\n\nif list(final_plan[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Final model plan target order does not match TARGETS.\"\n    )\n\nprint(\"Final OOF-backed model plan: PASS\")\nprint(\"Final target ordering: PASS\")\n\n# ----------------------------------------------------------------------\n# CHECKPOINT PATHS\n# ----------------------------------------------------------------------\n\ncheckpoint_paths = {\n    \"baseline_fold0\": os.path.join(\n        audit_dir,\n        \"cell37_fold0_best_model.pt\"\n    ),\n    \"baseline_fold1\": os.path.join(\n        audit_dir,\n        \"cell40_fold1_best_model.pt\"\n    ),\n    \"baseline_fold2\": os.path.join(\n        audit_dir,\n        \"cell44_fold2_best_model.pt\"\n    ),\n    \"finetuned_fold0\": os.path.join(\n        audit_dir,\n        \"cell52_fold0_finetuned_experiment.pt\"\n    ),\n    \"finetuned_fold1\": os.path.join(\n        audit_dir,\n        \"cell54_fold1_finetuned_experiment.pt\"\n    ),\n    \"finetuned_fold2\": os.path.join(\n        audit_dir,\n        \"cell56_fold2_finetuned_experiment.pt\"\n    )\n}\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT FILE VALIDATION\")\nprint(\"=\" * 70)\n\nfor name, path in checkpoint_paths.items():\n\n    if not os.path.exists(path):\n        raise RuntimeError(\n            f\"Missing checkpoint: {path}\"\n        )\n\n    if os.path.getsize(path) <= 0:\n        raise RuntimeError(\n            f\"Checkpoint is empty: {path}\"\n        )\n\n    print(f\"{name}: EXISTS\")\n\nprint(\"Six checkpoint files: PASS\")\n\n# ----------------------------------------------------------------------\n# CHECKPOINT STATE-DICT EXTRACTION\n# ----------------------------------------------------------------------\n\ndef extract_state_dict(checkpoint, checkpoint_name):\n    \"\"\"\n    Extract a PyTorch state_dict without assuming a single\n    checkpoint wrapper format.\n    \"\"\"\n\n    if isinstance(checkpoint, dict):\n\n        if \"state_dict\" in checkpoint:\n            state_dict = checkpoint[\"state_dict\"]\n\n        elif \"model_state_dict\" in checkpoint:\n            state_dict = checkpoint[\"model_state_dict\"]\n\n        elif \"model\" in checkpoint and isinstance(\n            checkpoint[\"model\"],\n            dict\n        ):\n            state_dict = checkpoint[\"model\"]\n\n        else:\n            # Some checkpoints are themselves state_dict dictionaries.\n            if all(\n                isinstance(k, str)\n                and isinstance(v, torch.Tensor)\n                for k, v in checkpoint.items()\n            ):\n                state_dict = checkpoint\n            else:\n                raise RuntimeError(\n                    f\"{checkpoint_name}: unable to identify state_dict.\"\n                )\n\n    else:\n        raise RuntimeError(\n            f\"{checkpoint_name}: unsupported checkpoint type \"\n            f\"{type(checkpoint)}.\"\n        )\n\n    if not isinstance(state_dict, dict):\n        raise RuntimeError(\n            f\"{checkpoint_name}: extracted state_dict is invalid.\"\n        )\n\n    if len(state_dict) == 0:\n        raise RuntimeError(\n            f\"{checkpoint_name}: state_dict is empty.\"\n        )\n\n    return state_dict\n\n# ----------------------------------------------------------------------\n# STATE-DICT COMPATIBILITY CHECK\n# ----------------------------------------------------------------------\n\ntemplate_state = model.state_dict()\n\nloaded_models = {}\n\nprint()\nprint(\"=\" * 70)\nprint(\"CHECKPOINT COMPATIBILITY\")\nprint(\"=\" * 70)\n\nfor checkpoint_name, checkpoint_path in checkpoint_paths.items():\n\n    checkpoint = torch.load(\n        checkpoint_path,\n        map_location=inference_device\n    )\n\n    state_dict = extract_state_dict(\n        checkpoint,\n        checkpoint_name\n    )\n\n    # Handle DataParallel-style \"module.\" prefixes if present.\n    normalized_state_dict = {}\n\n    for key, value in state_dict.items():\n\n        normalized_key = key\n\n        if normalized_key.startswith(\"module.\"):\n            normalized_key = normalized_key[len(\"module.\"):]\n\n        normalized_state_dict[normalized_key] = value\n\n    checkpoint_keys = set(normalized_state_dict.keys())\n    model_keys = set(template_state.keys())\n\n    missing_keys = sorted(\n        model_keys - checkpoint_keys\n    )\n\n    unexpected_keys = sorted(\n        checkpoint_keys - model_keys\n    )\n\n    if missing_keys or unexpected_keys:\n        raise RuntimeError(\n            f\"{checkpoint_name}: state_dict incompatible.\\n\"\n            f\"Missing keys: {missing_keys[:10]}\\n\"\n            f\"Unexpected keys: {unexpected_keys[:10]}\"\n        )\n\n    # Verify tensor shapes before loading.\n    shape_mismatches = []\n\n    for key in model_keys:\n\n        model_shape = tuple(\n            template_state[key].shape\n        )\n\n        checkpoint_shape = tuple(\n            normalized_state_dict[key].shape\n        )\n\n        if model_shape != checkpoint_shape:\n            shape_mismatches.append(\n                (\n                    key,\n                    model_shape,\n                    checkpoint_shape\n                )\n            )\n\n    if shape_mismatches:\n        raise RuntimeError(\n            f\"{checkpoint_name}: parameter shape mismatch: \"\n            + str(shape_mismatches[:5])\n        )\n\n    # Create an independent model copy.\n    candidate_model = copy.deepcopy(model)\n\n    candidate_model.load_state_dict(\n        normalized_state_dict,\n        strict=True\n    )\n\n    candidate_model = candidate_model.to(\n        inference_device\n    )\n\n    candidate_model.eval()\n\n    # Parameter validity.\n    invalid_parameters = False\n\n    for parameter in candidate_model.parameters():\n\n        if not torch.isfinite(\n            parameter\n        ).all():\n            invalid_parameters = True\n            break\n\n    if invalid_parameters:\n        raise RuntimeError(\n            f\"{checkpoint_name}: checkpoint contains \"\n            \"NaN/Inf parameters.\"\n        )\n\n    loaded_models[checkpoint_name] = candidate_model\n\n    print(\n        f\"{checkpoint_name}: \"\n        \"state_dict compatibility PASS | \"\n        \"model load PASS | \"\n        \"parameter validity PASS\"\n    )\n\nprint(\"All six checkpoints loaded successfully: PASS\")\n\n# ----------------------------------------------------------------------\n# MODEL ARCHITECTURE VERIFICATION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"ARCHITECTURE VERIFICATION\")\nprint(\"=\" * 70)\n\n# Use one known checkpoint model for a synthetic forward test.\n# This does NOT use any dataset or test image.\n\ntest_model = loaded_models[\"baseline_fold0\"]\n\nwith torch.no_grad():\n\n    synthetic_input = torch.zeros(\n        1,\n        21,\n        224,\n        224,\n        dtype=torch.float32,\n        device=inference_device\n    )\n\n    synthetic_output = test_model(\n        synthetic_input\n    )\n\nif not isinstance(\n    synthetic_output,\n    torch.Tensor\n):\n    raise RuntimeError(\n        \"Model forward output is not a torch.Tensor.\"\n    )\n\nif tuple(synthetic_output.shape) != (1, 12):\n    raise RuntimeError(\n        \"Unexpected model output shape: \"\n        + str(tuple(synthetic_output.shape))\n    )\n\nif not torch.isfinite(\n    synthetic_output\n).all():\n    raise RuntimeError(\n        \"Synthetic forward output contains NaN/Inf.\"\n    )\n\nprint(\"Synthetic input shape: (1, 21, 224, 224)\")\nprint(\"Synthetic output shape:\", tuple(synthetic_output.shape))\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Forward pass: PASS\")\nprint(\"Output validity: PASS\")\n\n# ----------------------------------------------------------------------\n# MODEL TYPE SUMMARY\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LOADED MODEL SUMMARY\")\nprint(\"=\" * 70)\n\nfor model_name, loaded_model in loaded_models.items():\n\n    total_parameters = sum(\n        parameter.numel()\n        for parameter in loaded_model.parameters()\n    )\n\n    print(\n        f\"{model_name}: \"\n        f\"parameters={total_parameters:,} | \"\n        f\"mode={'eval' if not loaded_model.training else 'train'}\"\n    )\n\n# ----------------------------------------------------------------------\n# FINAL PLAN / CHECKPOINT CONSISTENCY\n# ----------------------------------------------------------------------\n\nbaseline_targets = final_plan.loc[\n    final_plan[\"selected_model\"] == \"baseline\",\n    \"target\"\n].tolist()\n\nfinetuned_targets = final_plan.loc[\n    final_plan[\"selected_model\"] == \"fine_tuned\",\n    \"target\"\n].tolist()\n\nif len(baseline_targets) != 5:\n    raise RuntimeError(\n        f\"Expected 5 baseline-selected targets, found \"\n        f\"{len(baseline_targets)}.\"\n    )\n\nif len(finetuned_targets) != 7:\n    raise RuntimeError(\n        f\"Expected 7 fine-tuned-selected targets, found \"\n        f\"{len(finetuned_targets)}.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL MODEL PLAN CONSISTENCY\")\nprint(\"=\" * 70)\n\nprint(\"Baseline-selected targets: 5/12\")\nprint(\"Fine-tuned-selected targets: 7/12\")\nprint(\"All 12 targets assigned exactly once: PASS\")\nprint(\"OOF thresholds preserved: PASS\")\n\n# ----------------------------------------------------------------------\n# IMPORTANT: NO DATASET INFERENCE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"DATA SAFETY CHECK\")\nprint(\"=\" * 70)\n\nprint(\"No training performed.\")\nprint(\"No validation dataset inference performed.\")\nprint(\"No test data accessed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No checkpoint overwritten.\")\n\n# ----------------------------------------------------------------------\n# STORE LOADED MODELS FOR NEXT CELL\n# ----------------------------------------------------------------------\n\nfinal_inference_models = loaded_models\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 62 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Six validated checkpoints loaded: PASS\")\nprint(\"Baseline checkpoints: PASS\")\nprint(\"Fine-tuned checkpoints: PASS\")\nprint(\"State-dict compatibility: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Synthetic forward pass: PASS\")\nprint(\"Final OOF-backed target plan: PASS\")\nprint(\"Inference model objects prepared: PASS\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 62 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Final inference architecture is ready.\"\n)\n\nprint(\n    \"No dataset or test inference was performed.\"\n)\n\nprint(\n    \"Send the complete Cell 62 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:35:00.720145Z","iopub.execute_input":"2026-08-11T11:35:00.720949Z","iopub.status.idle":"2026-08-11T11:35:01.430436Z","shell.execute_reply.started":"2026-08-11T11:35:00.720911Z","shell.execute_reply":"2026-08-11T11:35:01.42938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 63 - FINAL 3-FOLD INFERENCE PLAN AUDIT\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 63 - FINAL 3-FOLD INFERENCE PLAN AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"final_inference_models\"\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# TARGET VALIDATION\n# ----------------------------------------------------------------------\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 targets, found {len(TARGETS)}.\"\n    )\n\nprint(f\"Targets: {len(TARGETS)}\")\n\n# ----------------------------------------------------------------------\n# LOADED MODEL VALIDATION\n# ----------------------------------------------------------------------\n\nexpected_model_names = {\n    \"baseline_fold0\",\n    \"baseline_fold1\",\n    \"baseline_fold2\",\n    \"finetuned_fold0\",\n    \"finetuned_fold1\",\n    \"finetuned_fold2\"\n}\n\nactual_model_names = set(\n    final_inference_models.keys()\n)\n\nif actual_model_names != expected_model_names:\n    raise RuntimeError(\n        \"Loaded model set mismatch.\\n\"\n        f\"Expected: {sorted(expected_model_names)}\\n\"\n        f\"Found: {sorted(actual_model_names)}\"\n    )\n\nprint(\"Six inference models: PASS\")\n\n# ----------------------------------------------------------------------\n# MODEL MODE VALIDATION\n# ----------------------------------------------------------------------\n\nfor model_name, loaded_model in final_inference_models.items():\n\n    if loaded_model.training:\n        raise RuntimeError(\n            f\"{model_name} is still in training mode.\"\n        )\n\nprint(\"All inference models in evaluation mode: PASS\")\n\n# ----------------------------------------------------------------------\n# LOAD FINAL OOF-BACKED PLAN\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nplan_path = os.path.join(\n    audit_dir,\n    \"cell61_final_oof_backed_model_plan.csv\"\n)\n\nif not os.path.exists(plan_path):\n    raise RuntimeError(\n        f\"Missing final OOF-backed plan: {plan_path}\"\n    )\n\nfinal_plan = pd.read_csv(plan_path)\n\nif len(final_plan) != 12:\n    raise RuntimeError(\n        \"Final plan must contain exactly 12 targets.\"\n    )\n\nrequired_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"selected_best_f1\"\n]\n\nmissing_columns = [\n    column\n    for column in required_columns\n    if column not in final_plan.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Final plan missing columns: \"\n        + \", \".join(missing_columns)\n    )\n\nif list(final_plan[\"target\"]) != list(TARGETS):\n    raise RuntimeError(\n        \"Final plan target order does not match TARGETS.\"\n    )\n\nprint(\"Final OOF-backed plan: PASS\")\nprint(\"Target ordering: PASS\")\n\n# ----------------------------------------------------------------------\n# VALIDATE MODEL ASSIGNMENTS\n# ----------------------------------------------------------------------\n\nvalid_model_assignments = {\n    \"baseline\",\n    \"fine_tuned\"\n}\n\ninvalid_assignments = [\n    value\n    for value in final_plan[\"selected_model\"]\n    if value not in valid_model_assignments\n]\n\nif invalid_assignments:\n    raise RuntimeError(\n        \"Invalid model assignments: \"\n        + \", \".join(\n            sorted(set(invalid_assignments))\n        )\n    )\n\nprint(\"Target model assignments: PASS\")\n\n# ----------------------------------------------------------------------\n# VALIDATE THRESHOLDS\n# ----------------------------------------------------------------------\n\nthresholds = pd.to_numeric(\n    final_plan[\"selected_threshold\"],\n    errors=\"coerce\"\n)\n\nif thresholds.isna().any():\n    raise RuntimeError(\n        \"Selected thresholds contain invalid values.\"\n    )\n\nif not thresholds.between(\n    0.10,\n    0.90\n).all():\n    raise RuntimeError(\n        \"Selected thresholds fall outside \"\n        \"the audited 0.10-0.90 range.\"\n    )\n\nprint(\"Threshold range: PASS\")\n\n# ----------------------------------------------------------------------\n# BUILD EXPLICIT THREE-FOLD MODEL MAP\n# ----------------------------------------------------------------------\n\nmodel_map = {}\n\nfor _, row in final_plan.iterrows():\n\n    target = row[\"target\"]\n    selected_model = row[\"selected_model\"]\n    threshold = float(\n        row[\"selected_threshold\"]\n    )\n\n    if selected_model == \"baseline\":\n\n        fold_models = [\n            \"baseline_fold0\",\n            \"baseline_fold1\",\n            \"baseline_fold2\"\n        ]\n\n    elif selected_model == \"fine_tuned\":\n\n        fold_models = [\n            \"finetuned_fold0\",\n            \"finetuned_fold1\",\n            \"finetuned_fold2\"\n        ]\n\n    else:\n        raise RuntimeError(\n            f\"Unsupported selected model: {selected_model}\"\n        )\n\n    model_map[target] = {\n        \"selected_model\": selected_model,\n        \"threshold\": threshold,\n        \"fold_models\": fold_models\n    }\n\n# ----------------------------------------------------------------------\n# VERIFY EVERY TARGET HAS THREE FOLD MODELS\n# ----------------------------------------------------------------------\n\nfor target in TARGETS:\n\n    if target not in model_map:\n        raise RuntimeError(\n            f\"Missing inference plan for target: {target}\"\n        )\n\n    fold_models = model_map[target][\"fold_models\"]\n\n    if len(fold_models) != 3:\n        raise RuntimeError(\n            f\"{target}: expected 3 fold models.\"\n        )\n\n    for fold_model in fold_models:\n\n        if fold_model not in final_inference_models:\n            raise RuntimeError(\n                f\"{target}: missing loaded model \"\n                f\"{fold_model}.\"\n            )\n\nprint(\"Every target has exactly three fold models: PASS\")\n\n# ----------------------------------------------------------------------\n# VERIFY BASELINE / FINE-TUNED COUNTS\n# ----------------------------------------------------------------------\n\nbaseline_targets = [\n    target\n    for target in TARGETS\n    if model_map[target][\"selected_model\"] == \"baseline\"\n]\n\nfinetuned_targets = [\n    target\n    for target in TARGETS\n    if model_map[target][\"selected_model\"] == \"fine_tuned\"\n]\n\nif len(baseline_targets) != 5:\n    raise RuntimeError(\n        f\"Expected 5 baseline targets, found \"\n        f\"{len(baseline_targets)}.\"\n    )\n\nif len(finetuned_targets) != 7:\n    raise RuntimeError(\n        f\"Expected 7 fine-tuned targets, found \"\n        f\"{len(finetuned_targets)}.\"\n    )\n\nprint(\"Baseline target count: 5\")\nprint(\"Fine-tuned target count: 7\")\nprint(\"Model-family assignment: PASS\")\n\n# ----------------------------------------------------------------------\n# PRINT FINAL INFERENCE PLAN\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL THREE-FOLD INFERENCE PLAN\")\nprint(\"=\" * 70)\n\nplan_rows = []\n\nfor target in TARGETS:\n\n    entry = model_map[target]\n\n    plan_rows.append({\n        \"target\": target,\n        \"selected_model\": entry[\"selected_model\"],\n        \"threshold\": entry[\"threshold\"],\n        \"fold0_model\": entry[\"fold_models\"][0],\n        \"fold1_model\": entry[\"fold_models\"][1],\n        \"fold2_model\": entry[\"fold_models\"][2]\n    })\n\ninference_plan_table = pd.DataFrame(\n    plan_rows\n)\n\nprint(\n    inference_plan_table.to_string(\n        index=False,\n        float_format=lambda x: f\"{x:.4f}\"\n    )\n)\n\n# ----------------------------------------------------------------------\n# SAVE INFERENCE PLAN\n# ----------------------------------------------------------------------\n\ninference_plan_path = os.path.join(\n    audit_dir,\n    \"cell63_final_three_fold_inference_plan.csv\"\n)\n\ninference_plan_table.to_csv(\n    inference_plan_path,\n    index=False\n)\n\nif not os.path.exists(inference_plan_path):\n    raise RuntimeError(\n        \"Failed to save final three-fold inference plan.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 63 VERIFICATION\")\nprint(\"=\" * 70)\n\nprint(\"12 targets covered: PASS\")\nprint(\"Three folds assigned per target: PASS\")\nprint(\"Baseline model family mapping: PASS\")\nprint(\"Fine-tuned model family mapping: PASS\")\nprint(\"OOF thresholds preserved: PASS\")\nprint(\"All six loaded checkpoints referenced: PASS\")\nprint(\"No training performed: PASS\")\nprint(\"No validation inference performed: PASS\")\nprint(\"No test data accessed: PASS\")\nprint(\"No checkpoint modified: PASS\")\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 63 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Saved: {inference_plan_path}\"\n)\n\nprint(\n    \"Final three-fold inference plan is ready.\"\n)\n\nprint(\n    \"Send the complete Cell 63 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:38:00.227447Z","iopub.execute_input":"2026-08-11T11:38:00.22777Z","iopub.status.idle":"2026-08-11T11:38:00.265398Z","shell.execute_reply.started":"2026-08-11T11:38:00.227741Z","shell.execute_reply":"2026-08-11T11:38:00.264231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 64R - TEST SOURCE RESTORATION\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 64R - TEST SOURCE RESTORATION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"KneeStudyDataset\",\n    \"test\",\n    \"test_series\",\n    \"test_slice_audit\",\n    \"TEST_SERIES_DIR\",\n    \"TEST_SERIES_CSV\",\n]\n\nmissing_objects = [\n    name for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# ESTABLISHED SCHEMA\n# ----------------------------------------------------------------------\n\nIDENTITY_COLUMN = \"StudyInstanceUID\"\n\nPRIMARY_SERIES_COLUMNS = [\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\",\n]\n\nTEST_SERIES_COLUMNS = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\",\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\",\n    \"Anatomical_Plane\",\n]\n\nprint()\nprint(\"=\" * 70)\nprint(\"ESTABLISHED TEST SCHEMA\")\nprint(\"=\" * 70)\n\nprint(\"Identity:\", IDENTITY_COLUMN)\n\nfor column in PRIMARY_SERIES_COLUMNS:\n    print(\"Primary series:\", column)\n\n# ----------------------------------------------------------------------\n# VALIDATE ACTUAL TEST SOURCE\n# ----------------------------------------------------------------------\n\nif not isinstance(test, pd.DataFrame):\n    raise RuntimeError(\n        \"Notebook object 'test' is not a DataFrame.\"\n    )\n\nif not isinstance(test_series, pd.DataFrame):\n    raise RuntimeError(\n        \"Notebook object 'test_series' is not a DataFrame.\"\n    )\n\nif not isinstance(test_slice_audit, pd.DataFrame):\n    raise RuntimeError(\n        \"Notebook object 'test_slice_audit' is not a DataFrame.\"\n    )\n\nmissing_test_columns = [\n    column\n    for column in [IDENTITY_COLUMN]\n    if column not in test.columns\n]\n\nif missing_test_columns:\n    raise RuntimeError(\n        \"test is missing required columns: \"\n        + \", \".join(missing_test_columns)\n    )\n\nmissing_series_columns = [\n    column\n    for column in TEST_SERIES_COLUMNS\n    if column not in test_series.columns\n]\n\nif missing_series_columns:\n    raise RuntimeError(\n        \"test_series is missing required columns: \"\n        + \", \".join(missing_series_columns)\n    )\n\nprint(\"test DataFrame: PASS\")\nprint(\"test_series DataFrame: PASS\")\nprint(\"test_slice_audit DataFrame: PASS\")\n\n# ----------------------------------------------------------------------\n# NORMALIZE IDENTIFIERS\n# ----------------------------------------------------------------------\n\ntest_ids = (\n    test[IDENTITY_COLUMN]\n    .astype(str)\n    .str.strip()\n)\n\ntest_series_work = test_series.copy()\n\ntest_series_work[IDENTITY_COLUMN] = (\n    test_series_work[IDENTITY_COLUMN]\n    .astype(str)\n    .str.strip()\n)\n\ntest_series_work[\"SeriesInstanceUID\"] = (\n    test_series_work[\"SeriesInstanceUID\"]\n    .astype(str)\n    .str.strip()\n)\n\n# Preserve exact test.csv ordering.\ntest_id_order = {\n    study_id: index\n    for index, study_id in enumerate(test_ids.tolist())\n}\n\ntest_id_set = set(test_ids.tolist())\n\nprint()\nprint(\"=\" * 70)\nprint(\"ACTUAL TEST SOURCE\")\nprint(\"=\" * 70)\n\nprint(\"TEST_SERIES_DIR:\", TEST_SERIES_DIR)\nprint(\"TEST_SERIES_CSV:\", TEST_SERIES_CSV)\nprint(\"Test studies from test:\", len(test_id_set))\nprint(\"Rows in test_series:\", len(test_series_work))\n\n# ----------------------------------------------------------------------\n# ISOLATE ONLY THE ACTUAL TEST STUDIES\n#\n# Do not require test_series to contain only test.csv studies globally.\n# The earlier 64S failure came from over-restrictive source\n# classification. We only need to safely extract the requested test IDs.\n# ----------------------------------------------------------------------\n\ntest_candidates = test_series_work[\n    test_series_work[IDENTITY_COLUMN].isin(test_id_set)\n].copy()\n\nif test_candidates.empty:\n    raise RuntimeError(\n        \"No rows from test_series matched the StudyInstanceUIDs \"\n        \"in test.\"\n    )\n\nmatched_test_ids = set(\n    test_candidates[IDENTITY_COLUMN]\n)\n\nmissing_test_studies = (\n    test_id_set - matched_test_ids\n)\n\nprint()\nprint(\"Matched test studies:\", len(matched_test_ids))\n\nif missing_test_studies:\n    print(\n        \"Studies from test.csv without test_series rows:\",\n        sorted(missing_test_studies),\n    )\n    raise RuntimeError(\n        \"Actual test studies are missing from test_series. \"\n        \"Do not fabricate test metadata.\"\n    )\n\nprint(\"Test study-to-series mapping: PASS\")\n\n# ----------------------------------------------------------------------\n# TEST SLICE-COUNT INFORMATION\n# ----------------------------------------------------------------------\n\nslice_audit = test_slice_audit.copy()\n\nslice_key_columns = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\",\n]\n\nmissing_slice_keys = [\n    column\n    for column in slice_key_columns\n    if column not in slice_audit.columns\n]\n\nif missing_slice_keys:\n    raise RuntimeError(\n        \"test_slice_audit is missing: \"\n        + \", \".join(missing_slice_keys)\n    )\n\nslice_audit[IDENTITY_COLUMN] = (\n    slice_audit[IDENTITY_COLUMN]\n    .astype(str)\n    .str.strip()\n)\n\nslice_audit[\"SeriesInstanceUID\"] = (\n    slice_audit[\"SeriesInstanceUID\"]\n    .astype(str)\n    .str.strip()\n)\n\nif \"slice_count\" in slice_audit.columns:\n\n    slice_counts = slice_audit[\n        slice_key_columns + [\"slice_count\"]\n    ].copy()\n\n    slice_counts[\"slice_count\"] = pd.to_numeric(\n        slice_counts[\"slice_count\"],\n        errors=\"coerce\",\n    )\n\n    # Avoid accidental many-to-many merge.\n    slice_counts = (\n        slice_counts\n        .drop_duplicates(\n            subset=slice_key_columns,\n            keep=\"first\",\n        )\n    )\n\n    test_candidates = test_candidates.merge(\n        slice_counts,\n        on=slice_key_columns,\n        how=\"left\",\n        validate=\"one_to_one\",\n    )\n\nelse:\n    # The test-series metadata itself is still sufficient for\n    # identifying the series. Do not fabricate slice counts.\n    test_candidates[\"slice_count\"] = np.nan\n\nprint(\"Test slice metadata handling: PASS\")\n\n# ----------------------------------------------------------------------\n# PLANE NORMALIZATION\n# ----------------------------------------------------------------------\n\ntest_candidates[\"Anatomical_Plane\"] = (\n    test_candidates[\"Anatomical_Plane\"]\n    .astype(str)\n    .str.strip()\n)\n\nexpected_planes = [\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\",\n]\n\n# ----------------------------------------------------------------------\n# SHOW ACTUAL CANDIDATES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"ACTUAL TEST SERIES CANDIDATES\")\nprint(\"=\" * 70)\n\ncandidate_display_columns = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\",\n    \"Anatomical_Plane\",\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\",\n    \"slice_count\",\n]\n\navailable_display_columns = [\n    column\n    for column in candidate_display_columns\n    if column in test_candidates.columns\n]\n\nprint(\n    test_candidates[\n        available_display_columns\n    ]\n    .sort_values(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n            \"SeriesInstanceUID\",\n        ]\n    )\n    .to_string(index=False)\n)\n\n# ----------------------------------------------------------------------\n# VALIDATE THREE-PLANE COVERAGE\n# ----------------------------------------------------------------------\n\ncoverage_rows = []\n\nfor study_id in test_ids.tolist():\n\n    study_candidates = test_candidates[\n        test_candidates[IDENTITY_COLUMN] == study_id\n    ]\n\n    observed_planes = set(\n        study_candidates[\"Anatomical_Plane\"]\n    )\n\n    missing_planes = [\n        plane\n        for plane in expected_planes\n        if plane not in observed_planes\n    ]\n\n    coverage_rows.append(\n        {\n            \"StudyInstanceUID\": study_id,\n            \"Sagittal\": int(\n                \"Sagittal\" in observed_planes\n            ),\n            \"Coronal\": int(\n                \"Coronal\" in observed_planes\n            ),\n            \"Axial\": int(\n                \"Axial\" in observed_planes\n            ),\n            \"missing_planes\": (\n                \",\".join(missing_planes)\n                if missing_planes\n                else \"\"\n            ),\n        }\n    )\n\ntest_plane_coverage = pd.DataFrame(\n    coverage_rows\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST THREE-PLANE COVERAGE\")\nprint(\"=\" * 70)\n\nprint(\n    test_plane_coverage.to_string(\n        index=False\n    )\n)\n\nif (\n    test_plane_coverage[\n        \"missing_planes\"\n    ].astype(str).str.len() > 0\n).any():\n\n    raise RuntimeError(\n        \"At least one test study is missing a required \"\n        \"primary plane. Do not fabricate a series.\"\n    )\n\nprint(\"Three-plane coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# PRIMARY SERIES SELECTION\n#\n# Use the same type of evidence already established in the notebook:\n#   1. Fluid-sensitive preference\n#   2. Fat-suppression preference\n#   3. Slice count when available\n#   4. SeriesInstanceUID only as a deterministic final ordering key\n#\n# No labels or fold information are involved.\n# ----------------------------------------------------------------------\n\ntest_candidates[\"_fluid_score\"] = pd.to_numeric(\n    test_candidates[\"Fluid_Sensitive\"],\n    errors=\"coerce\",\n).fillna(0)\n\ntest_candidates[\"_fat_score\"] = pd.to_numeric(\n    test_candidates[\"Fat_Suppression\"],\n    errors=\"coerce\",\n).fillna(0)\n\ntest_candidates[\"_slice_score\"] = pd.to_numeric(\n    test_candidates[\"slice_count\"],\n    errors=\"coerce\",\n).fillna(-1)\n\nselected_rows = []\n\nfor study_id in test_ids.tolist():\n\n    study_candidates = test_candidates[\n        test_candidates[IDENTITY_COLUMN] == study_id\n    ].copy()\n\n    for plane in expected_planes:\n\n        plane_candidates = study_candidates[\n            study_candidates[\"Anatomical_Plane\"] == plane\n        ].copy()\n\n        if plane_candidates.empty:\n            raise RuntimeError(\n                f\"No {plane} series found for study {study_id}.\"\n            )\n\n        # Deterministic ranking.\n        plane_candidates = plane_candidates.sort_values(\n            [\n                \"_fluid_score\",\n                \"_fat_score\",\n                \"_slice_score\",\n                \"SeriesInstanceUID\",\n            ],\n            ascending=[\n                False,\n                False,\n                False,\n                True,\n            ],\n        ).reset_index(drop=True)\n\n        selected_rows.append(\n            plane_candidates.iloc[0].copy()\n        )\n\nselected_test_series = pd.DataFrame(\n    selected_rows\n)\n\n# ----------------------------------------------------------------------\n# VERIFY EXACTLY ONE SELECTED SERIES PER STUDY / PLANE\n# ----------------------------------------------------------------------\n\nselection_counts = (\n    selected_test_series\n    .groupby(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n        ]\n    )\n    .size()\n)\n\nif not bool(\n    (selection_counts == 1).all()\n):\n    raise RuntimeError(\n        \"Primary-series selection produced duplicate \"\n        \"study/plane assignments.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"SELECTED PRIMARY TEST SERIES\")\nprint(\"=\" * 70)\n\nprint(\n    selected_test_series[\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n            \"SeriesInstanceUID\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"slice_count\",\n        ]\n    ].to_string(index=False)\n)\n\nprint(\"One selected series per study/plane: PASS\")\n\n# ----------------------------------------------------------------------\n# PIVOT INTO THE ESTABLISHED 4-COLUMN TEST METADATA\n# ----------------------------------------------------------------------\n\ntest_metadata = (\n    selected_test_series\n    .pivot(\n        index=\"StudyInstanceUID\",\n        columns=\"Anatomical_Plane\",\n        values=\"SeriesInstanceUID\",\n    )\n    .reset_index()\n)\n\nmissing_planes_after_pivot = [\n    plane\n    for plane in expected_planes\n    if plane not in test_metadata.columns\n]\n\nif missing_planes_after_pivot:\n    raise RuntimeError(\n        \"Pivot failed to produce: \"\n        + \", \".join(missing_planes_after_pivot)\n    )\n\ntest_metadata = test_metadata.rename(\n    columns={\n        \"Sagittal\":\n            \"Sagittal_SeriesInstanceUID\",\n        \"Coronal\":\n            \"Coronal_SeriesInstanceUID\",\n        \"Axial\":\n            \"Axial_SeriesInstanceUID\",\n    }\n)\n\ntest_metadata = test_metadata[\n    [\n        \"StudyInstanceUID\",\n        \"Sagittal_SeriesInstanceUID\",\n        \"Coronal_SeriesInstanceUID\",\n        \"Axial_SeriesInstanceUID\",\n    ]\n].copy()\n\n# Preserve test.csv order.\ntest_metadata[\"_order\"] = (\n    test_metadata[\"StudyInstanceUID\"]\n    .map(test_id_order)\n)\n\nif test_metadata[\"_order\"].isna().any():\n    raise RuntimeError(\n        \"Could not restore test.csv study ordering.\"\n    )\n\ntest_metadata = (\n    test_metadata\n    .sort_values(\"_order\")\n    .drop(columns=\"_order\")\n    .reset_index(drop=True)\n)\n\n# ----------------------------------------------------------------------\n# FINAL VALIDATION\n# ----------------------------------------------------------------------\n\nexpected_test_rows = len(test)\n\nif len(test_metadata) != expected_test_rows:\n    raise RuntimeError(\n        \"Final test metadata row count mismatch. \"\n        f\"Expected {expected_test_rows}, \"\n        f\"got {len(test_metadata)}.\"\n    )\n\nif (\n    test_metadata[\"StudyInstanceUID\"]\n    .duplicated()\n    .any()\n):\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID values in final test metadata.\"\n    )\n\nfor column in PRIMARY_SERIES_COLUMNS:\n\n    if test_metadata[column].isna().any():\n        raise RuntimeError(\n            f\"Missing values in {column}.\"\n        )\n\n    if (\n        test_metadata[column]\n        .astype(str)\n        .str.strip()\n        .eq(\"\")\n        .any()\n    ):\n        raise RuntimeError(\n            f\"Empty values in {column}.\"\n        )\n\nif set(\n    test_metadata[\"StudyInstanceUID\"]\n) != test_id_set:\n\n    raise RuntimeError(\n        \"Final test metadata does not contain exactly \"\n        \"the test.csv studies.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL RESTORED TEST METADATA\")\nprint(\"=\" * 70)\n\nprint(\n    test_metadata.to_string(index=False)\n)\n\nprint()\nprint(\"Test metadata shape:\", test_metadata.shape)\nprint(\"Expected test studies:\", expected_test_rows)\nprint(\"StudyInstanceUID coverage: PASS\")\nprint(\"Sagittal_SeriesInstanceUID: PASS\")\nprint(\"Coronal_SeriesInstanceUID: PASS\")\nprint(\"Axial_SeriesInstanceUID: PASS\")\nprint(\"One row per test study: PASS\")\n\n# ----------------------------------------------------------------------\n# SAVE\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nos.makedirs(\n    audit_dir,\n    exist_ok=True,\n)\n\ntest_metadata_path = os.path.join(\n    audit_dir,\n    \"cell64r_test_primary_series_metadata.csv\",\n)\n\ntest_metadata.to_csv(\n    test_metadata_path,\n    index=False,\n)\n\n# Expose the canonical object for the next cell.\ntest_metadata_df = test_metadata.copy()\n\n# ----------------------------------------------------------------------\n# VERDICT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 64R VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Actual test source recovered: PASS\")\nprint(\"Test studies:\", len(test_metadata_df))\nprint(\"Primary-series metadata: PASS\")\nprint(\"Established schema preserved: PASS\")\nprint(\"No labels fabricated: PASS\")\nprint(\"No fold assignments fabricated: PASS\")\nprint(\"No modeling table modified: PASS\")\nprint(\"No training performed: PASS\")\nprint(\"No checkpoint modified: PASS\")\nprint(\"No test inference performed: PASS\")\n\nprint()\nprint(\"Saved:\")\nprint(test_metadata_path)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 64R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Test metadata restoration completed.\"\n)\nprint(\n    \"Next step: construct the test Dataset using the \"\n    \"established KneeStudyDataset pipeline.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:52:27.190234Z","iopub.execute_input":"2026-08-11T11:52:27.190561Z","iopub.status.idle":"2026-08-11T11:52:27.294007Z","shell.execute_reply.started":"2026-08-11T11:52:27.190534Z","shell.execute_reply":"2026-08-11T11:52:27.293037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 65 - TEST SERIES SELECTION AUDIT\n# ======================================================================\n\nimport inspect\nimport pandas as pd\nimport numpy as np\n\nprint(\"=\" * 70)\nprint(\"CELL 65 - TEST SERIES SELECTION AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"KneeStudyDataset\",\n    \"test\",\n    \"test_series\",\n    \"sample_submission\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"TARGETS: AVAILABLE\")\nprint(\"KneeStudyDataset: AVAILABLE\")\nprint(\"test: AVAILABLE\")\nprint(\"test_series: AVAILABLE\")\nprint(\"sample_submission: AVAILABLE\")\n\n# ----------------------------------------------------------------------\n# TEST TABLE SHAPES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST SOURCE TABLES\")\nprint(\"=\" * 70)\n\nprint(\"test shape:\", test.shape)\nprint(\"test_series shape:\", test_series.shape)\nprint(\"sample_submission shape:\", sample_submission.shape)\n\n# ----------------------------------------------------------------------\n# EXACT TEST STUDY IDS\n# ----------------------------------------------------------------------\n\nif \"StudyInstanceUID\" not in test.columns:\n    raise RuntimeError(\n        \"test is missing StudyInstanceUID.\"\n    )\n\nif \"StudyInstanceUID\" not in test_series.columns:\n    raise RuntimeError(\n        \"test_series is missing StudyInstanceUID.\"\n    )\n\ntest_ids_from_test = set(\n    test[\"StudyInstanceUID\"].astype(str)\n)\n\ntest_ids_from_series = set(\n    test_series[\"StudyInstanceUID\"].astype(str)\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST STUDY COVERAGE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Studies in test.csv:\",\n    len(test_ids_from_test)\n)\n\nprint(\n    \"Studies in test_series.csv:\",\n    len(test_ids_from_series)\n)\n\nmissing_series_studies = (\n    test_ids_from_test\n    - test_ids_from_series\n)\n\nextra_series_studies = (\n    test_ids_from_series\n    - test_ids_from_test\n)\n\nprint(\n    \"Test studies missing from test_series:\",\n    len(missing_series_studies)\n)\n\nprint(\n    \"Extra studies in test_series:\",\n    len(extra_series_studies)\n)\n\nif missing_series_studies:\n    raise RuntimeError(\n        \"Some test.csv studies have no corresponding test_series rows.\"\n    )\n\nif extra_series_studies:\n    raise RuntimeError(\n        \"test_series contains studies not present in test.csv.\"\n    )\n\nprint(\"Test study coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# EXACT SERIES SCHEMA\n# ----------------------------------------------------------------------\n\nrequired_series_columns = [\n    \"StudyInstanceUID\",\n    \"SeriesInstanceUID\",\n    \"Fluid_Sensitive\",\n    \"Fat_Suppression\",\n    \"Anatomical_Plane\",\n]\n\nmissing_series_columns = [\n    column\n    for column in required_series_columns\n    if column not in test_series.columns\n]\n\nif missing_series_columns:\n    raise RuntimeError(\n        \"test_series missing columns: \"\n        + \", \".join(missing_series_columns)\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST SERIES SCHEMA\")\nprint(\"=\" * 70)\n\nfor column in required_series_columns:\n    print(\n        f\"{column}: PASS\"\n    )\n\n# ----------------------------------------------------------------------\n# SERIES ID VALIDATION\n# ----------------------------------------------------------------------\n\nif test_series[\"SeriesInstanceUID\"].isna().any():\n    raise RuntimeError(\n        \"test_series contains missing SeriesInstanceUID values.\"\n    )\n\nif test_series[\"SeriesInstanceUID\"].astype(str).str.strip().eq(\"\").any():\n    raise RuntimeError(\n        \"test_series contains empty SeriesInstanceUID values.\"\n    )\n\nif test_series[\"StudyInstanceUID\"].isna().any():\n    raise RuntimeError(\n        \"test_series contains missing StudyInstanceUID values.\"\n    )\n\nprint(\"SeriesInstanceUID validity: PASS\")\nprint(\"StudyInstanceUID validity: PASS\")\n\n# ----------------------------------------------------------------------\n# DUPLICATE ROW CHECK\n# ----------------------------------------------------------------------\n\nduplicate_series_rows = test_series[\n    test_series.duplicated(\n        subset=[\n            \"StudyInstanceUID\",\n            \"SeriesInstanceUID\",\n        ],\n        keep=False,\n    )\n]\n\nprint(\n    \"Duplicate StudyInstanceUID/SeriesInstanceUID rows:\",\n    len(duplicate_series_rows)\n)\n\nif len(duplicate_series_rows) > 0:\n    raise RuntimeError(\n        \"Duplicate study/series mappings found in test_series.\"\n    )\n\nprint(\"Study/series uniqueness: PASS\")\n\n# ----------------------------------------------------------------------\n# PLANE VALIDATION\n# ----------------------------------------------------------------------\n\nexpected_planes = {\n    \"Sagittal\",\n    \"Coronal\",\n    \"Axial\",\n}\n\nobserved_planes = set(\n    test_series[\"Anatomical_Plane\"]\n    .dropna()\n    .astype(str)\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"ANATOMICAL PLANE VALIDATION\")\nprint(\"=\" * 70)\n\nprint(\n    \"Observed planes:\",\n    sorted(observed_planes)\n)\n\nunexpected_planes = (\n    observed_planes\n    - expected_planes\n)\n\nif unexpected_planes:\n    raise RuntimeError(\n        \"Unexpected anatomical planes: \"\n        + \", \".join(sorted(unexpected_planes))\n    )\n\nmissing_planes = (\n    expected_planes\n    - observed_planes\n)\n\nif missing_planes:\n    raise RuntimeError(\n        \"Missing anatomical planes globally: \"\n        + \", \".join(sorted(missing_planes))\n    )\n\nprint(\"Anatomical planes: PASS\")\n\n# ----------------------------------------------------------------------\n# PER-STUDY PLANE COVERAGE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PER-STUDY PLANE COVERAGE\")\nprint(\"=\" * 70)\n\ncoverage_rows = []\n\nfor study_id, group in test_series.groupby(\n    \"StudyInstanceUID\",\n    sort=False,\n):\n\n    planes = set(\n        group[\"Anatomical_Plane\"]\n        .astype(str)\n    )\n\n    coverage_rows.append(\n        {\n            \"StudyInstanceUID\": study_id,\n            \"series_count\": len(group),\n            \"sagittal\": \"Sagittal\" in planes,\n            \"coronal\": \"Coronal\" in planes,\n            \"axial\": \"Axial\" in planes,\n        }\n    )\n\ntest_plane_coverage = pd.DataFrame(\n    coverage_rows\n)\n\nprint(test_plane_coverage.to_string(index=False))\n\ncoverage_ok = bool(\n    test_plane_coverage[\n        [\n            \"sagittal\",\n            \"coronal\",\n            \"axial\",\n        ]\n    ].all(axis=1).all()\n)\n\nif not coverage_ok:\n    raise RuntimeError(\n        \"At least one test study is missing a required anatomical plane.\"\n    )\n\nprint()\nprint(\"Three-plane coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# PRINT ALL TEST SERIES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"ALL TEST SERIES\")\nprint(\"=\" * 70)\n\nprint(\n    test_series[\n        required_series_columns\n    ].sort_values(\n        [\n            \"StudyInstanceUID\",\n            \"Anatomical_Plane\",\n            \"Fluid_Sensitive\",\n            \"Fat_Suppression\",\n            \"SeriesInstanceUID\",\n        ]\n    ).to_string(index=False)\n)\n\n# ----------------------------------------------------------------------\n# PER-PLANE CANDIDATE AUDIT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"PER-STUDY / PER-PLANE CANDIDATE AUDIT\")\nprint(\"=\" * 70)\n\ncandidate_rows = []\n\nfor study_id, study_group in test_series.groupby(\n    \"StudyInstanceUID\",\n    sort=False,\n):\n\n    for plane in [\n        \"Sagittal\",\n        \"Coronal\",\n        \"Axial\",\n    ]:\n\n        plane_group = study_group[\n            study_group[\"Anatomical_Plane\"].astype(str)\n            == plane\n        ].copy()\n\n        if plane_group.empty:\n            raise RuntimeError(\n                f\"Study {study_id} has no {plane} series.\"\n            )\n\n        for _, row in plane_group.iterrows():\n\n            candidate_rows.append(\n                {\n                    \"StudyInstanceUID\": study_id,\n                    \"Anatomical_Plane\": plane,\n                    \"SeriesInstanceUID\": row[\n                        \"SeriesInstanceUID\"\n                    ],\n                    \"Fluid_Sensitive\": row[\n                        \"Fluid_Sensitive\"\n                    ],\n                    \"Fat_Suppression\": row[\n                        \"Fat_Suppression\"\n                    ],\n                }\n            )\n\ntest_series_candidates = pd.DataFrame(\n    candidate_rows\n)\n\nprint(\n    test_series_candidates.to_string(\n        index=False\n    )\n)\n\n# ----------------------------------------------------------------------\n# EXISTING NOTEBOOK SELECTION OBJECTS\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"EXISTING SERIES-SELECTION OBJECTS\")\nprint(\"=\" * 70)\n\nfor object_name in [\n    \"selected_series\",\n    \"series_audit\",\n    \"test_summary\",\n    \"test_slice_audit\",\n    \"metadata_source\",\n    \"test_mapping_counts\",\n    \"test_bad_series_mapping\",\n]:\n\n    if object_name in globals():\n\n        value = globals()[object_name]\n\n        print()\n        print(\n            f\"{object_name}: \"\n            f\"type={type(value).__name__}\"\n        )\n\n        if isinstance(value, pd.DataFrame):\n\n            print(\n                \"shape:\",\n                value.shape\n            )\n\n            print(\n                \"columns:\",\n                list(value.columns)\n            )\n\n        elif isinstance(value, (list, tuple, set)):\n\n            print(\n                \"length:\",\n                len(value)\n            )\n\n# ----------------------------------------------------------------------\n# EXISTING FUNCTION INSPECTION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"EXISTING FUNCTION SIGNATURES\")\nprint(\"=\" * 70)\n\nfunctions_to_inspect = [\n    \"get_series_directory\",\n    \"get_series_dir\",\n    \"get_dicom_files\",\n    \"get_sorted_dicom_files\",\n    \"load_series\",\n    \"load_series_33\",\n    \"normalize_dicom_image\",\n    \"read_dicom_series\",\n]\n\nfor function_name in functions_to_inspect:\n\n    if function_name in globals():\n\n        function_object = globals()[function_name]\n\n        try:\n\n            signature = inspect.signature(\n                function_object\n            )\n\n            print(\n                f\"{function_name}{signature}\"\n            )\n\n        except Exception as exc:\n\n            print(\n                f\"{function_name}: \"\n                f\"signature unavailable \"\n                f\"({type(exc).__name__})\"\n            )\n\n# ----------------------------------------------------------------------\n# DATASET CONSTRUCTOR INSPECTION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"KNEESTUDYDATASET CONSTRUCTOR\")\nprint(\"=\" * 70)\n\ntry:\n\n    dataset_signature = inspect.signature(\n        KneeStudyDataset\n    )\n\n    print(\n        \"KneeStudyDataset:\",\n        dataset_signature\n    )\n\nexcept Exception as exc:\n\n    print(\n        \"Could not inspect KneeStudyDataset signature:\",\n        type(exc).__name__,\n        str(exc)\n    )\n\n# ----------------------------------------------------------------------\n# EXISTING TEST LOADER CANDIDATES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"EXISTING TEST LOADER CANDIDATES\")\nprint(\"=\" * 70)\n\nfor object_name in [\n    \"test_loader_candidates\",\n    \"available_test_loaders\",\n]:\n\n    if object_name in globals():\n\n        value = globals()[object_name]\n\n        print(\n            f\"{object_name}:\",\n            value\n        )\n\n# ----------------------------------------------------------------------\n# NO SELECTION / NO DATASET CREATION\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 65 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Test studies: {len(test_ids_from_test)}\"\n)\n\nprint(\n    f\"Test series rows: {len(test_series)}\"\n)\n\nprint(\n    \"Test study coverage: PASS\"\n)\n\nprint(\n    \"Test series schema: PASS\"\n)\n\nprint(\n    \"Three-plane coverage: PASS\"\n)\n\nprint(\n    \"Per-plane candidate audit: PASS\"\n)\n\nprint(\n    \"Existing selection objects inspected: PASS\"\n)\n\nprint(\n    \"Existing dataset constructor inspected: PASS\"\n)\n\nprint()\nprint(\n    \"NO TEST DATASET CREATED.\"\n)\n\nprint(\n    \"NO TEST DATALOADER CREATED.\"\n)\n\nprint(\n    \"NO TEST INFERENCE PERFORMED.\"\n)\n\nprint(\n    \"NO CHECKPOINT MODIFIED.\"\n)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 65 COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Send the complete Cell 65 output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:53:55.353582Z","iopub.execute_input":"2026-08-11T11:53:55.353962Z","iopub.status.idle":"2026-08-11T11:53:55.415534Z","shell.execute_reply.started":"2026-08-11T11:53:55.353926Z","shell.execute_reply":"2026-08-11T11:53:55.414635Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 66A - KNEESTUDYDATASET TEST-MODE COMPATIBILITY AUDIT\n# ======================================================================\n\nimport inspect\n\nprint(\"=\" * 70)\nprint(\"CELL 66A - KNEESTUDYDATASET TEST-MODE COMPATIBILITY AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"KneeStudyDataset\",\n    \"test_metadata\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# TEST METADATA VALIDATION\n# ----------------------------------------------------------------------\n\nrequired_columns = [\n    \"StudyInstanceUID\",\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\",\n]\n\nmissing_columns = [\n    column\n    for column in required_columns\n    if column not in test_metadata.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Test metadata missing required columns: \"\n        + \", \".join(missing_columns)\n    )\n\nif len(test_metadata) != 3:\n    raise RuntimeError(\n        f\"Expected 3 test studies, got {len(test_metadata)}.\"\n    )\n\nif test_metadata[\"StudyInstanceUID\"].duplicated().any():\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID values in test metadata.\"\n    )\n\nprint(\"Test metadata shape:\", test_metadata.shape)\nprint(\"Test metadata schema: PASS\")\nprint(\"Test study count: 3\")\nprint(\"One row per study: PASS\")\n\n# ----------------------------------------------------------------------\n# TARGET VALIDATION\n# ----------------------------------------------------------------------\n\nif len(TARGETS) != 12:\n    raise RuntimeError(\n        f\"Expected 12 TARGETS, got {len(TARGETS)}.\"\n    )\n\nprint(\"TARGETS: PASS\")\nprint(\"Target count:\", len(TARGETS))\n\n# ----------------------------------------------------------------------\n# CLASS SIGNATURE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"DATASET CLASS SIGNATURE\")\nprint(\"=\" * 70)\n\ntry:\n    class_signature = inspect.signature(\n        KneeStudyDataset\n    )\n    print(\n        \"KneeStudyDataset:\",\n        class_signature\n    )\nexcept Exception as exc:\n    raise RuntimeError(\n        \"Could not inspect KneeStudyDataset signature.\"\n    ) from exc\n\n# ----------------------------------------------------------------------\n# METHOD SIGNATURES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"DATASET METHOD SIGNATURES\")\nprint(\"=\" * 70)\n\nmethods_to_check = [\n    \"__init__\",\n    \"__len__\",\n    \"__getitem__\",\n]\n\nmethod_signatures = {}\n\nfor method_name in methods_to_check:\n\n    if not hasattr(KneeStudyDataset, method_name):\n        raise RuntimeError(\n            f\"KneeStudyDataset missing method: {method_name}\"\n        )\n\n    method_object = getattr(\n        KneeStudyDataset,\n        method_name\n    )\n\n    try:\n        signature = inspect.signature(\n            method_object\n        )\n\n        method_signatures[method_name] = signature\n\n        print(\n            f\"{method_name}{signature}\"\n        )\n\n    except Exception as exc:\n\n        print(\n            f\"{method_name}: \"\n            f\"signature unavailable \"\n            f\"({type(exc).__name__})\"\n        )\n\n# ----------------------------------------------------------------------\n# SOURCE INSPECTION\n#\n# We inspect the existing implementation rather than creating a second\n# Dataset class or guessing how test samples should be represented.\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"DATASET IMPLEMENTATION SOURCE\")\nprint(\"=\" * 70)\n\nsource_results = {}\n\nfor method_name in methods_to_check:\n\n    method_object = getattr(\n        KneeStudyDataset,\n        method_name\n    )\n\n    try:\n        source = inspect.getsource(\n            method_object\n        )\n\n        source_results[method_name] = source\n\n        print()\n        print(\n            f\"----- {method_name} -----\"\n        )\n        print(source)\n\n    except Exception as exc:\n\n        source_results[method_name] = None\n\n        print()\n        print(\n            f\"{method_name}: source unavailable \"\n            f\"({type(exc).__name__})\"\n        )\n\n# ----------------------------------------------------------------------\n# SOURCE AVAILABILITY CHECK\n# ----------------------------------------------------------------------\n\nunavailable_methods = [\n    name\n    for name, source in source_results.items()\n    if source is None\n]\n\nif unavailable_methods:\n    raise RuntimeError(\n        \"Could not inspect Dataset implementation for: \"\n        + \", \".join(unavailable_methods)\n        + \". \"\n        \"Do not guess the test Dataset behavior.\"\n    )\n\nprint()\nprint(\"Dataset implementation inspection: PASS\")\n\n# ----------------------------------------------------------------------\n# LABEL DEPENDENCY AUDIT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"LABEL DEPENDENCY AUDIT\")\nprint(\"=\" * 70)\n\nall_source = \"\\n\".join(\n    source_results.values()\n)\n\nlabel_related_terms = [\n    \"label\",\n    \"targets\",\n    \"TARGETS\",\n    \"dataframe[target\",\n    \"dataframe[targets\",\n]\n\ndetected_terms = []\n\nfor term in label_related_terms:\n\n    if term in all_source:\n        detected_terms.append(term)\n\nprint(\n    \"Label-related implementation references:\",\n    detected_terms\n)\n\n# ----------------------------------------------------------------------\n# TEST METADATA COLUMN ACCESS AUDIT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST METADATA ACCESS AUDIT\")\nprint(\"=\" * 70)\n\nfor column in required_columns:\n\n    if column in all_source:\n        print(\n            f\"{column}: referenced by Dataset implementation\"\n        )\n    else:\n        print(\n            f\"{column}: not directly referenced by Dataset implementation\"\n        )\n\n# ----------------------------------------------------------------------\n# IMPORTANT SAFETY STOP\n#\n# We do NOT instantiate the Dataset yet. The next cell will be written\n# from the actual implementation discovered here.\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 66A VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Test metadata available: PASS\")\nprint(\"Test metadata schema: PASS\")\nprint(\"12-target configuration: PASS\")\nprint(\"KneeStudyDataset signature: PASS\")\nprint(\"Dataset methods located: PASS\")\nprint(\"Dataset implementation inspected: PASS\")\nprint()\nprint(\"No test Dataset constructed.\")\nprint(\"No test DataLoader constructed.\")\nprint(\"No test inference performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No training performed.\")\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 66A COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Existing Dataset implementation has been audited.\"\n)\nprint(\n    \"Send the complete CELL 66A output before constructing \"\n    \"the test Dataset.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T11:57:09.285862Z","iopub.execute_input":"2026-08-11T11:57:09.287157Z","iopub.status.idle":"2026-08-11T11:57:09.311817Z","shell.execute_reply.started":"2026-08-11T11:57:09.287118Z","shell.execute_reply":"2026-08-11T11:57:09.310862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 66C-R - FINALIZE TEST DATALOADER\n# ======================================================================\n\nimport os\nimport torch\nfrom torch.utils.data import DataLoader\n\nprint(\"=\" * 70)\nprint(\"CELL 66C-R - FINALIZE TEST DATALOADER\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"test_metadata\",\n    \"test_dataset\",\n    \"TRAIN_SERIES_DIR\",\n    \"TEST_SERIES_DIR\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# TEST METADATA VALIDATION\n# ----------------------------------------------------------------------\n\nrequired_test_columns = [\n    \"StudyInstanceUID\",\n    \"Sagittal_SeriesInstanceUID\",\n    \"Coronal_SeriesInstanceUID\",\n    \"Axial_SeriesInstanceUID\",\n]\n\nmissing_columns = [\n    column\n    for column in required_test_columns\n    if column not in test_metadata.columns\n]\n\nif missing_columns:\n    raise RuntimeError(\n        \"Test metadata missing required columns: \"\n        + \", \".join(missing_columns)\n    )\n\nif len(test_metadata) != 3:\n    raise RuntimeError(\n        f\"Expected 3 test studies, got {len(test_metadata)}.\"\n    )\n\nprint(\"Test metadata schema: PASS\")\nprint(\"Test metadata shape:\", test_metadata.shape)\nprint(\"Test studies:\", len(test_metadata))\n\n# ----------------------------------------------------------------------\n# DATASET VALIDATION\n# ----------------------------------------------------------------------\n\nif len(test_dataset) != 3:\n    raise RuntimeError(\n        f\"Expected test Dataset length 3, got {len(test_dataset)}.\"\n    )\n\nif not isinstance(\n    test_dataset,\n    TestKneeStudyDataset,\n):\n    raise RuntimeError(\n        \"test_dataset is not the expected TestKneeStudyDataset.\"\n    )\n\nprint(\"Test Dataset: PASS\")\nprint(\"Test Dataset length:\", len(test_dataset))\nprint(\n    \"Test Dataset class:\",\n    type(test_dataset).__name__\n)\n\n# ----------------------------------------------------------------------\n# PATH SAFETY\n#\n# The original training directory must already be restored after\n# __getitem__ completes. We simply verify the current global value.\n# No local variable from __getitem__ is referenced here.\n# ----------------------------------------------------------------------\n\nif not os.path.isdir(TRAIN_SERIES_DIR):\n    raise RuntimeError(\n        \"TRAIN_SERIES_DIR is not a valid directory:\\n\"\n        + str(TRAIN_SERIES_DIR)\n    )\n\nif not os.path.isdir(TEST_SERIES_DIR):\n    raise RuntimeError(\n        \"TEST_SERIES_DIR is not a valid directory:\\n\"\n        + str(TEST_SERIES_DIR)\n    )\n\nprint(\"TRAIN_SERIES_DIR restored: PASS\")\nprint(\"TEST_SERIES_DIR available: PASS\")\n\n# ----------------------------------------------------------------------\n# CREATE TEST DATALOADER\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CREATING TEST DATALOADER\")\nprint(\"=\" * 70)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=2,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=False,\n)\n\nif len(test_loader) != 2:\n    raise RuntimeError(\n        f\"Expected 2 test batches, got {len(test_loader)}.\"\n    )\n\nprint(\"Test DataLoader: PASS\")\nprint(\"Test batches:\", len(test_loader))\n\n# ----------------------------------------------------------------------\n# DATALOADER SMOKE TEST\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST DATALOADER SMOKE TEST\")\nprint(\"=\" * 70)\n\nfirst_batch = next(iter(test_loader))\n\nrequired_batch_keys = [\n    \"image\",\n    \"label\",\n    \"study_id\",\n]\n\nmissing_batch_keys = [\n    key\n    for key in required_batch_keys\n    if key not in first_batch\n]\n\nif missing_batch_keys:\n    raise RuntimeError(\n        \"Test batch missing keys: \"\n        + \", \".join(missing_batch_keys)\n    )\n\nimages = first_batch[\"image\"]\nlabels = first_batch[\"label\"]\nstudy_ids = first_batch[\"study_id\"]\n\nif images.ndim != 4:\n    raise RuntimeError(\n        \"Unexpected image tensor dimensions: \"\n        + str(tuple(images.shape))\n    )\n\nif images.shape[1:] != (\n    21,\n    224,\n    224,\n):\n    raise RuntimeError(\n        \"Unexpected image batch shape: \"\n        + str(tuple(images.shape))\n    )\n\nif labels.ndim != 2:\n    raise RuntimeError(\n        \"Unexpected label tensor dimensions: \"\n        + str(tuple(labels.shape))\n    )\n\nif labels.shape[1] != 12:\n    raise RuntimeError(\n        \"Unexpected label batch shape: \"\n        + str(tuple(labels.shape))\n    )\n\nif not torch.isfinite(images).all():\n    raise RuntimeError(\n        \"NaN/Inf detected in test image batch.\"\n    )\n\nprint(\"Batch keys:\", list(first_batch.keys()))\nprint(\"Image batch shape:\", tuple(images.shape))\nprint(\"Label structural shape:\", tuple(labels.shape))\nprint(\"Study IDs:\", list(study_ids))\nprint(\"Image min:\", float(images.min()))\nprint(\"Image max:\", float(images.max()))\nprint(\"Image mean:\", float(images.mean()))\nprint(\"Image std:\", float(images.std()))\n\nprint(\"Batch structure: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target structural format: PASS\")\nprint(\"NaN/Inf check: PASS\")\n\n# ----------------------------------------------------------------------\n# COMPLETE TEST COVERAGE\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"TEST DATALOADER COVERAGE\")\nprint(\"=\" * 70)\n\nloader_study_ids = []\n\nfor batch in test_loader:\n\n    batch_images = batch[\"image\"]\n    batch_labels = batch[\"label\"]\n    batch_ids = batch[\"study_id\"]\n\n    if batch_images.shape[1:] != (\n        21,\n        224,\n        224,\n    ):\n        raise RuntimeError(\n            \"Invalid image shape encountered during full \"\n            \"test DataLoader iteration: \"\n            + str(tuple(batch_images.shape))\n        )\n\n    if batch_labels.shape[1] != 12:\n        raise RuntimeError(\n            \"Invalid label structural shape encountered \"\n            \"during full test DataLoader iteration: \"\n            + str(tuple(batch_labels.shape))\n        )\n\n    if not torch.isfinite(batch_images).all():\n        raise RuntimeError(\n            \"NaN/Inf detected during full test DataLoader iteration.\"\n        )\n\n    loader_study_ids.extend(\n        [\n            str(study_id)\n            for study_id in batch_ids\n        ]\n    )\n\nexpected_study_ids = (\n    test_metadata[\n        \"StudyInstanceUID\"\n    ]\n    .astype(str)\n    .tolist()\n)\n\nif len(loader_study_ids) != 3:\n    raise RuntimeError(\n        \"Expected exactly 3 studies from test DataLoader, \"\n        f\"got {len(loader_study_ids)}.\"\n    )\n\nif len(set(loader_study_ids)) != 3:\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID detected in test DataLoader.\"\n    )\n\nif set(loader_study_ids) != set(\n    expected_study_ids\n):\n    raise RuntimeError(\n        \"Test DataLoader study coverage does not match \"\n        \"test metadata.\"\n    )\n\nprint(\"Studies loaded:\", len(loader_study_ids))\nprint(\"Unique studies:\", len(set(loader_study_ids)))\nprint(\"Expected studies:\", len(expected_study_ids))\nprint(\"3/3 test studies covered exactly once: PASS\")\n\n# ----------------------------------------------------------------------\n# FINAL GLOBAL PATH SAFETY CHECK\n# ----------------------------------------------------------------------\n\nif TRAIN_SERIES_DIR != (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"train_series\"\n):\n    raise RuntimeError(\n        \"TRAIN_SERIES_DIR was not restored to the \"\n        \"established training-series directory.\"\n    )\n\nif TEST_SERIES_DIR != (\n    \"/kaggle/input/competitions/\"\n    \"rsna-knee-abnormality-detection/\"\n    \"test_series\"\n):\n    raise RuntimeError(\n        \"TEST_SERIES_DIR changed unexpectedly.\"\n    )\n\nprint(\"Established TRAIN_SERIES_DIR restored: PASS\")\nprint(\"Established TEST_SERIES_DIR preserved: PASS\")\n\n# ----------------------------------------------------------------------\n# FINAL VERDICT\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 66C-R VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Test metadata: PASS\")\nprint(\"Test Dataset: PASS\")\nprint(\"Test DataLoader: PASS\")\nprint(\"3-plane test loading: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target structural format: PASS\")\nprint(\"NaN/Inf validation: PASS\")\nprint(\"3/3 test studies covered exactly once: PASS\")\nprint(\"Training-series path restored: PASS\")\nprint(\"Test-series path preserved: PASS\")\nprint()\nprint(\"No labels fabricated.\")\nprint(\"No training performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No predictions generated.\")\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 66C-R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Test Dataset and DataLoader are ready for final inference.\"\n)\nprint(\n    \"Send the complete CELL 66C-R output before proceeding.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T12:01:01.521563Z","iopub.execute_input":"2026-08-11T12:01:01.521929Z","iopub.status.idle":"2026-08-11T12:01:06.152789Z","shell.execute_reply.started":"2026-08-11T12:01:01.52188Z","shell.execute_reply":"2026-08-11T12:01:06.151784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 67A-R - RESTORE FINAL INFERENCE OBJECTS\n# ======================================================================\n\nimport os\nimport copy\nimport torch\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 67A-R - RESTORE FINAL INFERENCE OBJECTS\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED NOTEBOOK OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"test_loader\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------------\n\nif \"device\" not in globals():\n    device = torch.device(\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\nprint(\"Device:\", device)\n\n# ----------------------------------------------------------------------\n# AUDIT DIRECTORY\n# ----------------------------------------------------------------------\n\naudit_dir = \"/kaggle/working/rsna_knee_audit\"\n\nif not os.path.isdir(audit_dir):\n    raise RuntimeError(\n        \"Audit directory not found: \"\n        + audit_dir\n    )\n\nprint(\"Audit directory: PASS\")\n\n# ----------------------------------------------------------------------\n# LOAD CELL 61 AUTHORITATIVE OOF-BACKED MODEL PLAN\n#\n# Cell 61 contains:\n# target\n# selected_model\n# selected_threshold\n# selected_best_f1\n# ----------------------------------------------------------------------\n\ncell61_path = os.path.join(\n    audit_dir,\n    \"cell61_final_oof_backed_model_plan.csv\",\n)\n\nif not os.path.isfile(cell61_path):\n    raise RuntimeError(\n        \"Cell-61 final OOF-backed model plan not found: \"\n        + cell61_path\n    )\n\ncell61_plan = pd.read_csv(\n    cell61_path\n)\n\nprint(\n    \"Cell-61 model plan:\",\n    cell61_plan.shape,\n)\n\nrequired_cell61_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n]\n\nmissing_cell61_columns = [\n    column\n    for column in required_cell61_columns\n    if column not in cell61_plan.columns\n]\n\nif missing_cell61_columns:\n    raise RuntimeError(\n        \"Cell-61 model plan missing required columns: \"\n        + \", \".join(missing_cell61_columns)\n    )\n\n# ----------------------------------------------------------------------\n# LOAD CELL 63 AUTHORITATIVE THREE-FOLD INFERENCE PLAN\n#\n# Cell 63 contains:\n# target\n# selected_model\n# threshold\n# fold0_model\n# fold1_model\n# fold2_model\n# ----------------------------------------------------------------------\n\ncell63_path = os.path.join(\n    audit_dir,\n    \"cell63_final_three_fold_inference_plan.csv\",\n)\n\nif not os.path.isfile(cell63_path):\n    raise RuntimeError(\n        \"Cell-63 final three-fold inference plan not found: \"\n        + cell63_path\n    )\n\ncell63_plan = pd.read_csv(\n    cell63_path\n)\n\nprint(\n    \"Cell-63 inference plan:\",\n    cell63_plan.shape,\n)\n\nrequired_cell63_columns = [\n    \"target\",\n    \"selected_model\",\n    \"threshold\",\n    \"fold0_model\",\n    \"fold1_model\",\n    \"fold2_model\",\n]\n\nmissing_cell63_columns = [\n    column\n    for column in required_cell63_columns\n    if column not in cell63_plan.columns\n]\n\nif missing_cell63_columns:\n    raise RuntimeError(\n        \"Cell-63 inference plan missing required columns: \"\n        + \", \".join(missing_cell63_columns)\n    )\n\n# ----------------------------------------------------------------------\n# BASIC PLAN VALIDATION\n# ----------------------------------------------------------------------\n\nif len(cell61_plan) != 12:\n    raise RuntimeError(\n        \"Cell-61 plan must contain exactly 12 targets.\"\n    )\n\nif len(cell63_plan) != 12:\n    raise RuntimeError(\n        \"Cell-63 plan must contain exactly 12 targets.\"\n    )\n\nexpected_targets = list(TARGETS)\n\nif cell61_plan[\n    \"target\"\n].tolist() != expected_targets:\n\n    raise RuntimeError(\n        \"Cell-61 target ordering does not match TARGETS.\"\n    )\n\nif cell63_plan[\n    \"target\"\n].tolist() != expected_targets:\n\n    raise RuntimeError(\n        \"Cell-63 target ordering does not match TARGETS.\"\n    )\n\nif cell61_plan[\n    \"target\"\n].duplicated().any():\n\n    raise RuntimeError(\n        \"Duplicate targets found in Cell-61 plan.\"\n    )\n\nif cell63_plan[\n    \"target\"\n].duplicated().any():\n\n    raise RuntimeError(\n        \"Duplicate targets found in Cell-63 plan.\"\n    )\n\nprint(\"Cell-61 target ordering: PASS\")\nprint(\"Cell-63 target ordering: PASS\")\nprint(\"12-target coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# VERIFY CELL 61 AND CELL 63 MODEL ASSIGNMENTS AGREE\n# ----------------------------------------------------------------------\n\ncell61_models = dict(\n    zip(\n        cell61_plan[\"target\"],\n        cell61_plan[\"selected_model\"],\n    )\n)\n\ncell63_models = dict(\n    zip(\n        cell63_plan[\"target\"],\n        cell63_plan[\"selected_model\"],\n    )\n)\n\nmodel_disagreements = []\n\nfor target in expected_targets:\n\n    if cell61_models[target] != cell63_models[target]:\n\n        model_disagreements.append(\n            (\n                target,\n                cell61_models[target],\n                cell63_models[target],\n            )\n        )\n\nif model_disagreements:\n    raise RuntimeError(\n        \"Cell-61 and Cell-63 selected-model disagreement: \"\n        + str(model_disagreements)\n    )\n\nprint(\n    \"Cell-61 vs Cell-63 selected-model consistency: PASS\"\n)\n\n# ----------------------------------------------------------------------\n# VERIFY CELL 61 THRESHOLDS AGAINST CELL 63 THRESHOLDS\n#\n# Cell 61 uses selected_threshold.\n# Cell 63 uses threshold.\n#\n# They must represent the same OOF-derived thresholds.\n# ----------------------------------------------------------------------\n\ncell61_thresholds = dict(\n    zip(\n        cell61_plan[\"target\"],\n        cell61_plan[\"selected_threshold\"],\n    )\n)\n\ncell63_thresholds = dict(\n    zip(\n        cell63_plan[\"target\"],\n        cell63_plan[\"threshold\"],\n    )\n)\n\nthreshold_disagreements = []\n\nfor target in expected_targets:\n\n    threshold61 = float(\n        cell61_thresholds[target]\n    )\n\n    threshold63 = float(\n        cell63_thresholds[target]\n    )\n\n    if abs(\n        threshold61 - threshold63\n    ) > 1e-12:\n\n        threshold_disagreements.append(\n            (\n                target,\n                threshold61,\n                threshold63,\n            )\n        )\n\nif threshold_disagreements:\n    raise RuntimeError(\n        \"Cell-61 and Cell-63 threshold disagreement: \"\n        + str(threshold_disagreements)\n    )\n\nprint(\n    \"Cell-61 vs Cell-63 threshold consistency: PASS\"\n)\n\n# ----------------------------------------------------------------------\n# CONSTRUCT THE AUTHORITATIVE FINAL INFERENCE PLAN\n#\n# Threshold comes from Cell 61.\n# Fold mappings come from Cell 63.\n#\n# No threshold is recomputed.\n# No model assignment is changed.\n# ----------------------------------------------------------------------\n\nfinal_inference_plan = pd.DataFrame(\n    {\n        \"target\": expected_targets,\n        \"selected_model\": [\n            cell61_models[target]\n            for target in expected_targets\n        ],\n        \"selected_threshold\": [\n            float(cell61_thresholds[target])\n            for target in expected_targets\n        ],\n        \"fold0_model\": [\n            cell63_plan.loc[\n                cell63_plan[\"target\"] == target,\n                \"fold0_model\",\n            ].iloc[0]\n            for target in expected_targets\n        ],\n        \"fold1_model\": [\n            cell63_plan.loc[\n                cell63_plan[\"target\"] == target,\n                \"fold1_model\",\n            ].iloc[0]\n            for target in expected_targets\n        ],\n        \"fold2_model\": [\n            cell63_plan.loc[\n                cell63_plan[\"target\"] == target,\n                \"fold2_model\",\n            ].iloc[0]\n            for target in expected_targets\n        ],\n    }\n)\n\n# ----------------------------------------------------------------------\n# FINAL PLAN SCHEMA VALIDATION\n# ----------------------------------------------------------------------\n\nrequired_final_plan_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"fold0_model\",\n    \"fold1_model\",\n    \"fold2_model\",\n]\n\nmissing_final_plan_columns = [\n    column\n    for column in required_final_plan_columns\n    if column not in final_inference_plan.columns\n]\n\nif missing_final_plan_columns:\n    raise RuntimeError(\n        \"Final inference plan missing columns: \"\n        + \", \".join(missing_final_plan_columns)\n    )\n\nif len(final_inference_plan) != 12:\n    raise RuntimeError(\n        \"Final inference plan must contain 12 targets.\"\n    )\n\nif final_inference_plan[\n    \"target\"\n].tolist() != expected_targets:\n\n    raise RuntimeError(\n        \"Final inference plan target ordering invalid.\"\n    )\n\n# ----------------------------------------------------------------------\n# VERIFY SELECTED MODEL VALUES\n# ----------------------------------------------------------------------\n\nallowed_models = {\n    \"baseline\",\n    \"fine_tuned\",\n}\n\ninvalid_models = sorted(\n    set(\n        final_inference_plan[\n            \"selected_model\"\n        ]\n    ) - allowed_models\n)\n\nif invalid_models:\n    raise RuntimeError(\n        \"Invalid selected model values: \"\n        + \", \".join(invalid_models)\n    )\n\n# ----------------------------------------------------------------------\n# VERIFY THRESHOLDS\n# ----------------------------------------------------------------------\n\nif (\n    final_inference_plan[\n        \"selected_threshold\"\n    ].isna().any()\n):\n\n    raise RuntimeError(\n        \"Missing selected threshold.\"\n    )\n\nif (\n    (final_inference_plan[\"selected_threshold\"] < 0.0)\n    | (final_inference_plan[\"selected_threshold\"] > 1.0)\n).any():\n\n    raise RuntimeError(\n        \"Selected threshold outside [0, 1].\"\n    )\n\nprint(\"Final inference plan schema: PASS\")\nprint(\"Selected model values: PASS\")\nprint(\"Selected thresholds: PASS\")\nprint()\nprint(final_inference_plan.to_string(index=False))\n\n# ----------------------------------------------------------------------\n# CHECKPOINT PATHS\n#\n# These are the exact validated checkpoints from the completed\n# baseline and controlled fine-tuning experiments.\n# ----------------------------------------------------------------------\n\ncheckpoint_paths = {\n    \"baseline_fold0\": os.path.join(\n        audit_dir,\n        \"cell37_fold0_best_model.pt\",\n    ),\n    \"baseline_fold1\": os.path.join(\n        audit_dir,\n        \"cell40_fold1_best_model.pt\",\n    ),\n    \"baseline_fold2\": os.path.join(\n        audit_dir,\n        \"cell44_fold2_best_model.pt\",\n    ),\n    \"finetuned_fold0\": os.path.join(\n        audit_dir,\n        \"cell52_fold0_finetuned_experiment.pt\",\n    ),\n    \"finetuned_fold1\": os.path.join(\n        audit_dir,\n        \"cell54_fold1_finetuned_experiment.pt\",\n    ),\n    \"finetuned_fold2\": os.path.join(\n        audit_dir,\n        \"cell56_fold2_finetuned_experiment.pt\",\n    ),\n}\n\nfor model_name, checkpoint_path in checkpoint_paths.items():\n\n    if not os.path.isfile(\n        checkpoint_path\n    ):\n        raise RuntimeError(\n            f\"Missing validated checkpoint: \"\n            f\"{checkpoint_path}\"\n        )\n\nprint(\"Six checkpoint files: PASS\")\n\n# ----------------------------------------------------------------------\n# FIND EXISTING VALIDATED MODEL ARCHITECTURE\n#\n# Cell 62 established:\n# total parameters = 11,367,372\n# input = 21 channels\n# output = 12 targets\n#\n# Do not invent a new constructor here.\n# ----------------------------------------------------------------------\n\narchitecture_model = None\narchitecture_name = None\n\npreferred_model_names = [\n    \"test_model\",\n    \"model\",\n]\n\nfor name in preferred_model_names:\n\n    if name not in globals():\n        continue\n\n    candidate = globals()[name]\n\n    if not isinstance(\n        candidate,\n        torch.nn.Module,\n    ):\n        continue\n\n    parameter_count = sum(\n        parameter.numel()\n        for parameter in candidate.parameters()\n    )\n\n    if parameter_count == 11_367_372:\n\n        architecture_model = candidate\n        architecture_name = name\n        break\n\nif architecture_model is None:\n\n    for name, candidate in globals().items():\n\n        if not isinstance(\n            candidate,\n            torch.nn.Module,\n        ):\n            continue\n\n        try:\n            parameter_count = sum(\n                parameter.numel()\n                for parameter in candidate.parameters()\n            )\n        except Exception:\n            continue\n\n        if parameter_count == 11_367_372:\n\n            architecture_model = candidate\n            architecture_name = name\n            break\n\nif architecture_model is None:\n    raise RuntimeError(\n        \"No existing validated 11,367,372-parameter \"\n        \"model architecture is available. \"\n        \"Do not invent a model constructor.\"\n    )\n\nprint(\n    \"Existing validated architecture:\",\n    architecture_name,\n)\n\n# ----------------------------------------------------------------------\n# ARCHITECTURE VALIDATION\n# ----------------------------------------------------------------------\n\narchitecture_model = architecture_model.to(\n    device\n)\n\narchitecture_model.eval()\n\nsynthetic_input = torch.zeros(\n    1,\n    21,\n    224,\n    224,\n    dtype=torch.float32,\n    device=device,\n)\n\nwith torch.no_grad():\n\n    synthetic_output = architecture_model(\n        synthetic_input\n    )\n\nif synthetic_output.shape != (\n    1,\n    12,\n):\n\n    raise RuntimeError(\n        \"Architecture output shape is \"\n        + str(tuple(synthetic_output.shape))\n    )\n\nif not torch.isfinite(\n    synthetic_output\n).all():\n\n    raise RuntimeError(\n        \"Architecture synthetic output contains NaN/Inf.\"\n    )\n\nprint(\"Architecture: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Synthetic forward pass: PASS\")\n\n# ----------------------------------------------------------------------\n# CHECKPOINT STATE-DICT EXTRACTION\n# ----------------------------------------------------------------------\n\ndef get_state_dict_from_validated_checkpoint(\n    checkpoint,\n    model_name,\n):\n\n    if not isinstance(\n        checkpoint,\n        dict,\n    ):\n\n        raise RuntimeError(\n            f\"Unexpected checkpoint type for {model_name}: \"\n            f\"{type(checkpoint)}\"\n        )\n\n    if \"state_dict\" in checkpoint:\n\n        state_dict = checkpoint[\n            \"state_dict\"\n        ]\n\n    elif \"model_state_dict\" in checkpoint:\n\n        state_dict = checkpoint[\n            \"model_state_dict\"\n        ]\n\n    else:\n\n        # A raw PyTorch state_dict is itself a dictionary\n        # whose values are tensors.\n        if checkpoint and all(\n            isinstance(value, torch.Tensor)\n            for value in checkpoint.values()\n        ):\n\n            state_dict = checkpoint\n\n        else:\n\n            raise RuntimeError(\n                f\"Checkpoint {model_name} does not contain \"\n                \"a recognizable state_dict.\"\n            )\n\n    if not state_dict:\n        raise RuntimeError(\n            f\"Empty state_dict for {model_name}.\"\n        )\n\n    return state_dict\n\n# ----------------------------------------------------------------------\n# LOAD ALL SIX VALIDATED CHECKPOINTS\n# ----------------------------------------------------------------------\n\nmodel_registry = {}\n\nfor model_name, checkpoint_path in checkpoint_paths.items():\n\n    checkpoint = torch.load(\n        checkpoint_path,\n        map_location=device,\n    )\n\n    state_dict = (\n        get_state_dict_from_validated_checkpoint(\n            checkpoint,\n            model_name,\n        )\n    )\n\n    restored_model = copy.deepcopy(\n        architecture_model\n    )\n\n    try:\n\n        restored_model.load_state_dict(\n            state_dict,\n            strict=True,\n        )\n\n    except RuntimeError as exc:\n\n        raise RuntimeError(\n            f\"State-dict compatibility failed for \"\n            f\"{model_name}.\"\n        ) from exc\n\n    restored_model = restored_model.to(\n        device\n    )\n\n    restored_model.eval()\n\n    # Parameter validity\n    for parameter in restored_model.parameters():\n\n        if not torch.isfinite(\n            parameter\n        ).all():\n\n            raise RuntimeError(\n                f\"NaN/Inf parameter found in \"\n                f\"{model_name}.\"\n            )\n\n    # Forward validation\n    with torch.no_grad():\n\n        output = restored_model(\n            synthetic_input\n        )\n\n    if output.shape != (\n        1,\n        12,\n    ):\n\n        raise RuntimeError(\n            f\"Invalid output shape for \"\n            f\"{model_name}: \"\n            f\"{tuple(output.shape)}\"\n        )\n\n    if not torch.isfinite(\n        output\n    ).all():\n\n        raise RuntimeError(\n            f\"NaN/Inf output for \"\n            f\"{model_name}.\"\n        )\n\n    model_registry[\n        model_name\n    ] = restored_model\n\n    print(\n        f\"{model_name}: LOAD PASS\"\n    )\n\n# ----------------------------------------------------------------------\n# CREATE THE EXACT OBJECT NAMES REQUIRED BY CELL 67\n# ----------------------------------------------------------------------\n\nbaseline_fold0 = model_registry[\n    \"baseline_fold0\"\n]\n\nbaseline_fold1 = model_registry[\n    \"baseline_fold1\"\n]\n\nbaseline_fold2 = model_registry[\n    \"baseline_fold2\"\n]\n\nfinetuned_fold0 = model_registry[\n    \"finetuned_fold0\"\n]\n\nfinetuned_fold1 = model_registry[\n    \"finetuned_fold1\"\n]\n\nfinetuned_fold2 = model_registry[\n    \"finetuned_fold2\"\n]\n\n# ----------------------------------------------------------------------\n# VERIFY ALL MODELS ARE IN EVALUATION MODE\n# ----------------------------------------------------------------------\n\nfor model_name, model_object in model_registry.items():\n\n    if model_object.training:\n\n        raise RuntimeError(\n            f\"{model_name} is not in evaluation mode.\"\n        )\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 67A-R VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Cell-61 model plan: PASS\")\nprint(\"Cell-63 fold mapping plan: PASS\")\nprint(\"Model-selection consistency: PASS\")\nprint(\"Threshold consistency: PASS\")\nprint(\"Final combined inference plan: PASS\")\nprint(\"12 targets covered: PASS\")\nprint(\"Existing architecture: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Six validated checkpoints: PASS\")\nprint(\"All six checkpoint loads: PASS\")\nprint(\"All six models in evaluation mode: PASS\")\n\nprint()\nprint(\"No training performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No test inference performed.\")\n\nprint(\"=\" * 70)\nprint(\"CELL 67A-R COMPLETE\")\nprint(\"=\" * 70)\n\nprint(\n    \"Final inference objects are restored and validated.\"\n)\nprint(\n    \"Existing test_loader preserved.\"\n)\nprint(\n    \"Proceed to Cell 67 final test inference.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T12:04:45.162672Z","iopub.execute_input":"2026-08-11T12:04:45.163074Z","iopub.status.idle":"2026-08-11T12:04:46.296501Z","shell.execute_reply.started":"2026-08-11T12:04:45.163041Z","shell.execute_reply":"2026-08-11T12:04:46.295491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 67 - FINAL THREE-FOLD TEST INFERENCE\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\n\nprint(\"=\" * 70)\nprint(\"CELL 67 - FINAL THREE-FOLD TEST INFERENCE\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED OBJECTS\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"test_loader\",\n    \"final_inference_plan\",\n    \"baseline_fold0\",\n    \"baseline_fold1\",\n    \"baseline_fold2\",\n    \"finetuned_fold0\",\n    \"finetuned_fold1\",\n    \"finetuned_fold2\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n        + \". Do not continue.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\n\n# ----------------------------------------------------------------------\n# DEVICE\n# ----------------------------------------------------------------------\n\nif \"device\" not in globals():\n    device = torch.device(\n        \"cuda\"\n        if torch.cuda.is_available()\n        else \"cpu\"\n    )\n\nprint(\"Device:\", device)\n\n# ----------------------------------------------------------------------\n# VALIDATE FINAL PLAN\n# ----------------------------------------------------------------------\n\nrequired_plan_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"fold0_model\",\n    \"fold1_model\",\n    \"fold2_model\",\n]\n\nmissing_plan_columns = [\n    column\n    for column in required_plan_columns\n    if column not in final_inference_plan.columns\n]\n\nif missing_plan_columns:\n    raise RuntimeError(\n        \"Final inference plan missing columns: \"\n        + \", \".join(missing_plan_columns)\n    )\n\nif len(final_inference_plan) != len(TARGETS):\n    raise RuntimeError(\n        \"Final inference plan does not contain exactly \"\n        f\"{len(TARGETS)} targets.\"\n    )\n\nif final_inference_plan[\n    \"target\"\n].tolist() != list(TARGETS):\n\n    raise RuntimeError(\n        \"Final inference plan target ordering does not \"\n        \"match TARGETS.\"\n    )\n\nprint(\"Final inference plan: PASS\")\nprint(\"Target ordering: PASS\")\nprint(\"12-target configuration: PASS\")\n\n# ----------------------------------------------------------------------\n# MODEL REGISTRY\n# ----------------------------------------------------------------------\n\nmodel_registry = {\n    \"baseline_fold0\": baseline_fold0,\n    \"baseline_fold1\": baseline_fold1,\n    \"baseline_fold2\": baseline_fold2,\n    \"finetuned_fold0\": finetuned_fold0,\n    \"finetuned_fold1\": finetuned_fold1,\n    \"finetuned_fold2\": finetuned_fold2,\n}\n\n# ----------------------------------------------------------------------\n# MODEL VALIDATION\n# ----------------------------------------------------------------------\n\nfor model_name, model_object in model_registry.items():\n\n    if not isinstance(\n        model_object,\n        torch.nn.Module,\n    ):\n        raise RuntimeError(\n            f\"{model_name} is not a PyTorch model.\"\n        )\n\n    if model_object.training:\n        raise RuntimeError(\n            f\"{model_name} is not in evaluation mode.\"\n        )\n\nprint(\"Six inference models: PASS\")\nprint(\"Evaluation mode: PASS\")\n\n# ----------------------------------------------------------------------\n# VALIDATE TARGET MODEL MAPPING\n# ----------------------------------------------------------------------\n\nfor _, row in final_inference_plan.iterrows():\n\n    target = row[\"target\"]\n\n    for fold_column in [\n        \"fold0_model\",\n        \"fold1_model\",\n        \"fold2_model\",\n    ]:\n\n        model_name = row[fold_column]\n\n        if model_name not in model_registry:\n            raise RuntimeError(\n                f\"Invalid model '{model_name}' for \"\n                f\"{target} in {fold_column}.\"\n            )\n\n    selected_model = row[\"selected_model\"]\n\n    if selected_model not in [\n        \"baseline\",\n        \"fine_tuned\",\n    ]:\n        raise RuntimeError(\n            f\"Invalid selected_model '{selected_model}' \"\n            f\"for target {target}.\"\n        )\n\n    threshold = float(\n        row[\"selected_threshold\"]\n    )\n\n    if not 0.0 <= threshold <= 1.0:\n        raise RuntimeError(\n            f\"Invalid threshold {threshold} \"\n            f\"for target {target}.\"\n        )\n\nprint(\"Target model assignments: PASS\")\nprint(\"Threshold validation: PASS\")\n\n# ----------------------------------------------------------------------\n# TEST LOADER VALIDATION\n# ----------------------------------------------------------------------\n\nexpected_test_count = 3\n\nif len(test_loader.dataset) != expected_test_count:\n    raise RuntimeError(\n        \"Unexpected test dataset size: \"\n        + str(len(test_loader.dataset))\n    )\n\nprint(\n    \"Test Dataset size:\",\n    len(test_loader.dataset),\n)\n\nprint(\"Test DataLoader: PASS\")\n\n# ----------------------------------------------------------------------\n# EXPECTED TEST STUDY IDS\n#\n# Prefer the already-established test metadata object.\n# Do not guess IDs.\n# ----------------------------------------------------------------------\n\nif \"test_metadata\" in globals():\n\n    test_metadata_ids = [\n        str(value)\n        for value in test_metadata[\n            \"StudyInstanceUID\"\n        ].tolist()\n    ]\n\nelif \"test\" in globals() and (\n    \"StudyInstanceUID\" in test.columns\n):\n\n    test_metadata_ids = [\n        str(value)\n        for value in test[\n            \"StudyInstanceUID\"\n        ].tolist()\n    ]\n\nelif \"test_ids\" in globals():\n\n    test_metadata_ids = sorted(\n        str(value)\n        for value in test_ids\n    )\n\nelse:\n    raise RuntimeError(\n        \"No established test StudyInstanceUID source \"\n        \"is available.\"\n    )\n\nif len(test_metadata_ids) != expected_test_count:\n    raise RuntimeError(\n        \"Expected exactly 3 test StudyInstanceUIDs, \"\n        f\"found {len(test_metadata_ids)}.\"\n    )\n\nif len(set(test_metadata_ids)) != expected_test_count:\n    raise RuntimeError(\n        \"Duplicate test StudyInstanceUIDs detected.\"\n    )\n\nprint(\"Expected test studies:\", len(test_metadata_ids))\nprint(\"Test StudyInstanceUID source: PASS\")\n\n# ----------------------------------------------------------------------\n# INFERENCE STORAGE\n#\n# Store probabilities by study ID and target.\n# ----------------------------------------------------------------------\n\nprobability_rows = []\n\nseen_studies = []\n\n# ----------------------------------------------------------------------\n# FINAL TEST INFERENCE\n# ----------------------------------------------------------------------\n\nfor batch_index, batch in enumerate(test_loader):\n\n    if not isinstance(batch, dict):\n        raise RuntimeError(\n            f\"Unexpected test batch type at batch \"\n            f\"{batch_index}: {type(batch)}\"\n        )\n\n    required_batch_keys = [\n        \"image\",\n        \"study_id\",\n    ]\n\n    missing_batch_keys = [\n        key\n        for key in required_batch_keys\n        if key not in batch\n    ]\n\n    if missing_batch_keys:\n        raise RuntimeError(\n            f\"Test batch {batch_index} missing keys: \"\n            + \", \".join(missing_batch_keys)\n        )\n\n    images = batch[\"image\"]\n\n    study_ids = [\n        str(value)\n        for value in batch[\"study_id\"]\n    ]\n\n    if images.ndim != 4:\n        raise RuntimeError(\n            f\"Unexpected image batch shape: \"\n            f\"{tuple(images.shape)}\"\n        )\n\n    if images.shape[1:] != (\n        21,\n        224,\n        224,\n    ):\n        raise RuntimeError(\n            \"Unexpected test image shape: \"\n            + str(tuple(images.shape))\n        )\n\n    if not torch.isfinite(\n        images\n    ).all():\n\n        raise RuntimeError(\n            f\"NaN/Inf detected in test batch \"\n            f\"{batch_index}.\"\n        )\n\n    if len(study_ids) != images.shape[0]:\n        raise RuntimeError(\n            f\"Study ID count does not match batch size \"\n            f\"for batch {batch_index}.\"\n        )\n\n    images = images.to(\n        device,\n        non_blocking=True,\n    )\n\n    # --------------------------------------------------------------\n    # Run every model required by the final target plan.\n    # This computes each model once per test batch.\n    # --------------------------------------------------------------\n\n    batch_model_probabilities = {}\n\n    with torch.no_grad():\n\n        for model_name, model_object in model_registry.items():\n\n            model_object.eval()\n\n            logits = model_object(\n                images\n            )\n\n            if logits.shape != (\n                images.shape[0],\n                len(TARGETS),\n            ):\n\n                raise RuntimeError(\n                    f\"{model_name} produced unexpected \"\n                    f\"output shape: {tuple(logits.shape)}\"\n                )\n\n            if not torch.isfinite(\n                logits\n            ).all():\n\n                raise RuntimeError(\n                    f\"NaN/Inf logits from \"\n                    f\"{model_name}.\"\n                )\n\n            probabilities = torch.sigmoid(\n                logits\n            )\n\n            if not torch.isfinite(\n                probabilities\n            ).all():\n\n                raise RuntimeError(\n                    f\"NaN/Inf probabilities from \"\n                    f\"{model_name}.\"\n                )\n\n            batch_model_probabilities[\n                model_name\n            ] = (\n                probabilities\n                .detach()\n                .cpu()\n                .numpy()\n            )\n\n    # --------------------------------------------------------------\n    # TARGET-WISE THREE-FOLD ENSEMBLE\n    # --------------------------------------------------------------\n\n    for sample_index, study_id in enumerate(\n        study_ids\n    ):\n\n        row = {\n            \"StudyInstanceUID\": study_id\n        }\n\n        for target_index, target in enumerate(\n            TARGETS\n        ):\n\n            plan_row = final_inference_plan[\n                final_inference_plan[\"target\"]\n                == target\n            ]\n\n            if len(plan_row) != 1:\n                raise RuntimeError(\n                    f\"Expected exactly one plan row \"\n                    f\"for target {target}.\"\n                )\n\n            plan_row = plan_row.iloc[0]\n\n            fold_model_names = [\n                plan_row[\"fold0_model\"],\n                plan_row[\"fold1_model\"],\n                plan_row[\"fold2_model\"],\n            ]\n\n            fold_probabilities = np.asarray(\n                [\n                    batch_model_probabilities[\n                        model_name\n                    ][\n                        sample_index,\n                        target_index\n                    ]\n                    for model_name in fold_model_names\n                ],\n                dtype=np.float64,\n            )\n\n            if fold_probabilities.shape != (\n                3,\n            ):\n                raise RuntimeError(\n                    f\"Expected three fold probabilities \"\n                    f\"for {target}.\"\n                )\n\n            if not np.isfinite(\n                fold_probabilities\n            ).all():\n\n                raise RuntimeError(\n                    f\"Invalid fold probabilities \"\n                    f\"for {target}.\"\n                )\n\n            ensemble_probability = float(\n                np.mean(\n                    fold_probabilities\n                )\n            )\n\n            if not np.isfinite(\n                ensemble_probability\n            ):\n                raise RuntimeError(\n                    f\"Invalid ensemble probability \"\n                    f\"for {target}.\"\n                )\n\n            row[\n                target\n            ] = ensemble_probability\n\n        probability_rows.append(row)\n        seen_studies.append(study_id)\n\n    print(\n        f\"Batch {batch_index + 1}: \"\n        f\"{len(study_ids)} studies inferred\"\n    )\n\n# ----------------------------------------------------------------------\n# BUILD PROBABILITY TABLE\n# ----------------------------------------------------------------------\n\nif len(probability_rows) != expected_test_count:\n    raise RuntimeError(\n        \"Unexpected number of test prediction rows: \"\n        + str(len(probability_rows))\n    )\n\ntest_probability_df = pd.DataFrame(\n    probability_rows\n)\n\nexpected_probability_columns = [\n    \"StudyInstanceUID\"\n] + list(TARGETS)\n\nif test_probability_df.columns.tolist() != (\n    expected_probability_columns\n):\n\n    raise RuntimeError(\n        \"Test probability column ordering is invalid.\"\n    )\n\n# ----------------------------------------------------------------------\n# STUDY COVERAGE VALIDATION\n# ----------------------------------------------------------------------\n\nif len(test_probability_df) != expected_test_count:\n    raise RuntimeError(\n        \"Test probability table must contain exactly 3 rows.\"\n    )\n\nactual_ids = [\n    str(value)\n    for value in test_probability_df[\n        \"StudyInstanceUID\"\n    ].tolist()\n]\n\nif len(set(actual_ids)) != expected_test_count:\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUIDs in predictions.\"\n    )\n\nif set(actual_ids) != set(\n    test_metadata_ids\n):\n\n    raise RuntimeError(\n        \"Predicted StudyInstanceUIDs do not exactly \"\n        \"match the established test studies.\"\n    )\n\nprint(\"Test study coverage: PASS\")\nprint(\"3/3 test studies predicted exactly once: PASS\")\n\n# ----------------------------------------------------------------------\n# PROBABILITY VALIDATION\n# ----------------------------------------------------------------------\n\nprobability_values = test_probability_df[\n    TARGETS\n].to_numpy(\n    dtype=np.float64\n)\n\nif not np.isfinite(\n    probability_values\n).all():\n\n    raise RuntimeError(\n        \"Final probabilities contain NaN/Inf.\"\n    )\n\nif (\n    (probability_values < 0.0)\n    | (probability_values > 1.0)\n).any():\n\n    raise RuntimeError(\n        \"Final probabilities outside [0, 1].\"\n    )\n\nprint(\"Probability validity: PASS\")\n\n# ----------------------------------------------------------------------\n# APPLY OOF-DERIVED TARGET-SPECIFIC THRESHOLDS\n# ----------------------------------------------------------------------\n\ntest_submission_df = pd.DataFrame()\n\ntest_submission_df[\n    \"StudyInstanceUID\"\n] = test_probability_df[\n    \"StudyInstanceUID\"\n]\n\nfor target in TARGETS:\n\n    plan_row = final_inference_plan[\n        final_inference_plan[\"target\"]\n        == target\n    ]\n\n    if len(plan_row) != 1:\n        raise RuntimeError(\n            f\"Missing unique plan row for {target}.\"\n        )\n\n    threshold = float(\n        plan_row.iloc[0][\n            \"selected_threshold\"\n        ]\n    )\n\n    probabilities = test_probability_df[\n        target\n    ].to_numpy(\n        dtype=np.float64\n    )\n\n    test_submission_df[\n        target\n    ] = (\n        probabilities >= threshold\n    ).astype(\n        np.int64\n    )\n\n# ----------------------------------------------------------------------\n# SUBMISSION SCHEMA\n# ----------------------------------------------------------------------\n\nexpected_submission_columns = [\n    \"StudyInstanceUID\"\n] + list(TARGETS)\n\nif test_submission_df.columns.tolist() != (\n    expected_submission_columns\n):\n\n    raise RuntimeError(\n        \"Final submission column ordering does not \"\n        \"match the established 13-column schema.\"\n    )\n\nif test_submission_df.shape != (\n    expected_test_count,\n    1 + len(TARGETS),\n):\n\n    raise RuntimeError(\n        \"Unexpected final submission shape: \"\n        + str(test_submission_df.shape)\n    )\n\n# ----------------------------------------------------------------------\n# BINARY OUTPUT VALIDATION\n# ----------------------------------------------------------------------\n\nprediction_values = test_submission_df[\n    TARGETS\n].to_numpy()\n\nunique_prediction_values = set(\n    np.unique(\n        prediction_values\n    ).tolist()\n)\n\nif not unique_prediction_values.issubset(\n    {0, 1}\n):\n\n    raise RuntimeError(\n        \"Final predictions are not binary.\"\n    )\n\nprint(\"Threshold application: PASS\")\nprint(\"Binary prediction validation: PASS\")\nprint(\"Submission schema: PASS\")\n\n# ----------------------------------------------------------------------\n# DISPLAY FINAL PROBABILITIES\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL TEST PROBABILITIES\")\nprint(\"=\" * 70)\n\nprint(\n    test_probability_df.to_string(\n        index=False\n    )\n)\n\n# ----------------------------------------------------------------------\n# DISPLAY FINAL BINARY PREDICTIONS\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL TEST PREDICTIONS\")\nprint(\"=\" * 70)\n\nprint(\n    test_submission_df.to_string(\n        index=False\n    )\n)\n\n# ----------------------------------------------------------------------\n# SAVE FINAL PROBABILITY TABLE\n# ----------------------------------------------------------------------\n\nprobability_path = os.path.join(\n    audit_dir,\n    \"cell67_final_test_probabilities.csv\",\n)\n\ntest_probability_df.to_csv(\n    probability_path,\n    index=False,\n)\n\nif not os.path.isfile(\n    probability_path\n):\n\n    raise RuntimeError(\n        \"Final probability file was not created.\"\n    )\n\n# ----------------------------------------------------------------------\n# SAVE FINAL SUBMISSION\n# ----------------------------------------------------------------------\n\nsubmission_path = os.path.join(\n    audit_dir,\n    \"cell67_final_submission.csv\",\n)\n\ntest_submission_df.to_csv(\n    submission_path,\n    index=False,\n)\n\nif not os.path.isfile(\n    submission_path\n):\n\n    raise RuntimeError(\n        \"Final submission file was not created.\"\n    )\n\n# ----------------------------------------------------------------------\n# FINAL VALIDATION\n# ----------------------------------------------------------------------\n\nsaved_submission = pd.read_csv(\n    submission_path\n)\n\nif saved_submission.shape != (\n    expected_test_count,\n    1 + len(TARGETS),\n):\n\n    raise RuntimeError(\n        \"Saved submission shape validation failed.\"\n    )\n\nif saved_submission.columns.tolist() != (\n    expected_submission_columns\n):\n\n    raise RuntimeError(\n        \"Saved submission column validation failed.\"\n    )\n\nsaved_ids = [\n    str(value)\n    for value in saved_submission[\n        \"StudyInstanceUID\"\n    ].tolist()\n]\n\nif set(saved_ids) != set(\n    test_metadata_ids\n):\n\n    raise RuntimeError(\n        \"Saved submission StudyInstanceUID coverage failed.\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 67 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Existing test DataLoader used: PASS\")\nprint(\"3 test studies loaded: PASS\")\nprint(\"21-channel input: PASS\")\nprint(\"12-target output: PASS\")\nprint(\"Three-fold inference: PASS\")\nprint(\"Target-wise model selection preserved: PASS\")\nprint(\"OOF-derived thresholds preserved: PASS\")\nprint(\"Probability validity: PASS\")\nprint(\"Study coverage: PASS\")\nprint(\"Submission schema: PASS\")\nprint(\"Binary predictions: PASS\")\nprint(\"Final submission saved: PASS\")\nprint(\"No training performed: PASS\")\nprint(\"No checkpoint modified: PASS\")\nprint(\"No validation data used: PASS\")\n\nprint()\nprint(\"Probability file:\")\nprint(probability_path)\n\nprint()\nprint(\"Final submission:\")\nprint(submission_path)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 67 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T12:07:10.603403Z","iopub.execute_input":"2026-08-11T12:07:10.603766Z","iopub.status.idle":"2026-08-11T12:07:13.82911Z","shell.execute_reply.started":"2026-08-11T12:07:10.603726Z","shell.execute_reply":"2026-08-11T12:07:13.827999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 68R - FINAL SUBMISSION INTEGRITY AUDIT\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 68R - FINAL SUBMISSION INTEGRITY AUDIT\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# Established paths\n# ----------------------------------------------------------------------\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nFINAL_SUBMISSION_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell67_final_submission.csv\"\n)\n\nPROBABILITY_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell67_final_test_probabilities.csv\"\n)\n\nCELL61_PLAN_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell61_final_oof_backed_model_plan.csv\"\n)\n\nCELL63_PLAN_PATH = os.path.join(\n    AUDIT_DIR,\n    \"cell63_final_three_fold_inference_plan.csv\"\n)\n\n# ----------------------------------------------------------------------\n# Established target ordering\n# ----------------------------------------------------------------------\n\nEXPECTED_TARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n]\n\nEXPECTED_COLUMNS = [\n    \"StudyInstanceUID\"\n] + EXPECTED_TARGETS\n\n# ----------------------------------------------------------------------\n# Required notebook objects\n# ----------------------------------------------------------------------\n\nrequired_objects = [\n    \"TARGETS\",\n    \"test\",\n    \"sample_submission\",\n]\n\nmissing_objects = [\n    name\n    for name in required_objects\n    if name not in globals()\n]\n\nif missing_objects:\n    raise RuntimeError(\n        \"Missing notebook objects: \"\n        + \", \".join(missing_objects)\n    )\n\nif list(TARGETS) != EXPECTED_TARGETS:\n    raise RuntimeError(\n        \"Notebook TARGETS ordering does not match the established \"\n        \"12-target configuration.\"\n    )\n\nprint(\"Required notebook objects: PASS\")\nprint(\"Target count:\", len(TARGETS))\nprint(\"Target ordering: PASS\")\n\n# ----------------------------------------------------------------------\n# Required files\n# ----------------------------------------------------------------------\n\nrequired_files = {\n    \"Final submission\": FINAL_SUBMISSION_PATH,\n    \"Final probability file\": PROBABILITY_PATH,\n    \"Cell-61 model plan\": CELL61_PLAN_PATH,\n    \"Cell-63 inference plan\": CELL63_PLAN_PATH,\n}\n\nfor name, path in required_files.items():\n    if not os.path.isfile(path):\n        raise RuntimeError(\n            f\"{name} not found: {path}\"\n        )\n\n    print(f\"{name}: EXISTS\")\n\n# ----------------------------------------------------------------------\n# Load final submission\n# ----------------------------------------------------------------------\n\nsubmission = pd.read_csv(\n    FINAL_SUBMISSION_PATH\n)\n\nprint()\nprint(\"Final submission shape:\", submission.shape)\n\n# ----------------------------------------------------------------------\n# Submission schema\n# ----------------------------------------------------------------------\n\nif list(submission.columns) != EXPECTED_COLUMNS:\n    raise RuntimeError(\n        \"Final submission schema mismatch.\\n\"\n        f\"Expected: {EXPECTED_COLUMNS}\\n\"\n        f\"Found:    {list(submission.columns)}\"\n    )\n\nif list(sample_submission.columns) != EXPECTED_COLUMNS:\n    raise RuntimeError(\n        \"sample_submission schema mismatch.\"\n    )\n\nprint(\"Submission schema: PASS\")\nprint(\"Column ordering: PASS\")\nprint(\"Column count:\", len(submission.columns))\nprint(\"sample_submission schema: PASS\")\n\n# ----------------------------------------------------------------------\n# Row count\n# ----------------------------------------------------------------------\n\nif len(submission) != len(test):\n    raise RuntimeError(\n        f\"Submission row count {len(submission)} does not match \"\n        f\"test study count {len(test)}.\"\n    )\n\nif len(submission) != 3:\n    raise RuntimeError(\n        f\"Expected 3 test studies, found {len(submission)}.\"\n    )\n\nprint(\"Submission row count:\", len(submission))\nprint(\"Expected test studies: 3\")\nprint(\"Row count: PASS\")\n\n# ----------------------------------------------------------------------\n# StudyInstanceUID validation\n# ----------------------------------------------------------------------\n\nsubmission_ids = submission[\n    \"StudyInstanceUID\"\n].astype(str)\n\ntest_ids = test[\n    \"StudyInstanceUID\"\n].astype(str)\n\nif submission_ids.isna().any():\n    raise RuntimeError(\n        \"Submission contains missing StudyInstanceUID values.\"\n    )\n\nif submission_ids.duplicated().any():\n    raise RuntimeError(\n        \"Submission contains duplicate StudyInstanceUID values.\"\n    )\n\nif set(submission_ids) != set(test_ids):\n    raise RuntimeError(\n        \"Submission StudyInstanceUID set does not exactly match \"\n        \"the recovered test studies.\"\n    )\n\nprint(\"StudyInstanceUID completeness: PASS\")\nprint(\"StudyInstanceUID uniqueness: PASS\")\nprint(\"Test-study coverage: PASS\")\nprint(\"3/3 expected test studies present exactly once: PASS\")\n\n# ----------------------------------------------------------------------\n# Target validation\n# ----------------------------------------------------------------------\n\nfor target in EXPECTED_TARGETS:\n\n    if submission[target].isna().any():\n        raise RuntimeError(\n            f\"Target '{target}' contains missing values.\"\n        )\n\n    numeric_values = pd.to_numeric(\n        submission[target],\n        errors=\"raise\"\n    )\n\n    unique_values = set(\n        numeric_values.tolist()\n    )\n\n    if not unique_values.issubset({0, 1}):\n        raise RuntimeError(\n            f\"Target '{target}' contains non-binary values: \"\n            f\"{sorted(unique_values)}\"\n        )\n\n    if not np.isfinite(\n        numeric_values.to_numpy(dtype=np.float32)\n    ).all():\n        raise RuntimeError(\n            f\"Target '{target}' contains NaN or Inf values.\"\n        )\n\nprint(\"All target columns present: PASS\")\nprint(\"No missing target values: PASS\")\nprint(\"Binary target values: PASS\")\nprint(\"Target numeric validity: PASS\")\nprint(\"NaN/Inf validation: PASS\")\n\n# ----------------------------------------------------------------------\n# Sample submission study coverage\n# ----------------------------------------------------------------------\n\nsample_ids = (\n    sample_submission[\"StudyInstanceUID\"]\n    .astype(str)\n)\n\nif set(sample_ids) != set(submission_ids):\n    raise RuntimeError(\n        \"Submission studies do not match sample_submission studies.\"\n    )\n\nprint(\"Sample-submission study coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# Probability file validation\n# ----------------------------------------------------------------------\n\nprobabilities = pd.read_csv(\n    PROBABILITY_PATH\n)\n\nexpected_probability_columns = [\n    \"StudyInstanceUID\"\n] + EXPECTED_TARGETS\n\nif list(probabilities.columns) != expected_probability_columns:\n    raise RuntimeError(\n        \"Probability file schema mismatch.\"\n    )\n\nif len(probabilities) != 3:\n    raise RuntimeError(\n        \"Probability file must contain exactly 3 studies.\"\n    )\n\nprobability_ids = (\n    probabilities[\"StudyInstanceUID\"]\n    .astype(str)\n)\n\nif probability_ids.duplicated().any():\n    raise RuntimeError(\n        \"Probability file contains duplicate StudyInstanceUID values.\"\n    )\n\nif set(probability_ids) != set(test_ids):\n    raise RuntimeError(\n        \"Probability file study coverage does not match test studies.\"\n    )\n\nprobability_matrix = probabilities[\n    EXPECTED_TARGETS\n].to_numpy(dtype=np.float32)\n\nif not np.isfinite(probability_matrix).all():\n    raise RuntimeError(\n        \"Probability file contains NaN or Inf values.\"\n    )\n\nif (\n    (probability_matrix < 0.0)\n    | (probability_matrix > 1.0)\n).any():\n    raise RuntimeError(\n        \"Probability file contains values outside [0, 1].\"\n    )\n\nprint(\"Probability file schema: PASS\")\nprint(\"Probability validity: PASS\")\nprint(\"Probability range [0,1]: PASS\")\nprint(\"Probability study coverage: PASS\")\n\n# ======================================================================\n# CRITICAL PLAN VALIDATION\n#\n# Cell 61 = authoritative OOF-backed model + threshold selection\n# Cell 63 = authoritative three-fold model mapping\n#\n# Cell 63 DOES NOT contain selected_threshold.\n# Therefore thresholds MUST come from Cell 61.\n# ======================================================================\n\ncell61_plan = pd.read_csv(\n    CELL61_PLAN_PATH\n)\n\ncell63_plan = pd.read_csv(\n    CELL63_PLAN_PATH\n)\n\n# ----------------------------------------------------------------------\n# Cell 61 schema\n# ----------------------------------------------------------------------\n\nrequired_cell61_columns = [\n    \"target\",\n    \"selected_model\",\n    \"selected_threshold\",\n    \"selected_best_f1\",\n]\n\nmissing_cell61_columns = [\n    column\n    for column in required_cell61_columns\n    if column not in cell61_plan.columns\n]\n\nif missing_cell61_columns:\n    raise RuntimeError(\n        \"Cell-61 model plan missing required columns: \"\n        + \", \".join(missing_cell61_columns)\n    )\n\nif len(cell61_plan) != 12:\n    raise RuntimeError(\n        f\"Cell-61 model plan must contain 12 targets, \"\n        f\"found {len(cell61_plan)}.\"\n    )\n\nif list(cell61_plan[\"target\"]) != EXPECTED_TARGETS:\n    raise RuntimeError(\n        \"Cell-61 target ordering does not match TARGETS.\"\n    )\n\nprint(\"Cell-61 model plan schema: PASS\")\nprint(\"Cell-61 target ordering: PASS\")\nprint(\"Cell-61 target coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# Cell 63 schema\n#\n# Deliberately do NOT require selected_threshold here.\n# ----------------------------------------------------------------------\n\nrequired_cell63_columns = [\n    \"target\",\n    \"selected_model\",\n    \"fold0_model\",\n    \"fold1_model\",\n    \"fold2_model\",\n]\n\nmissing_cell63_columns = [\n    column\n    for column in required_cell63_columns\n    if column not in cell63_plan.columns\n]\n\nif missing_cell63_columns:\n    raise RuntimeError(\n        \"Cell-63 inference plan missing required columns: \"\n        + \", \".join(missing_cell63_columns)\n    )\n\nif len(cell63_plan) != 12:\n    raise RuntimeError(\n        f\"Cell-63 inference plan must contain 12 targets, \"\n        f\"found {len(cell63_plan)}.\"\n    )\n\nif list(cell63_plan[\"target\"]) != EXPECTED_TARGETS:\n    raise RuntimeError(\n        \"Cell-63 target ordering does not match TARGETS.\"\n    )\n\nprint(\"Cell-63 model-mapping schema: PASS\")\nprint(\"Cell-63 target ordering: PASS\")\nprint(\"Cell-63 target coverage: PASS\")\n\n# ----------------------------------------------------------------------\n# Cell 61 vs Cell 63 model-selection consistency\n# ----------------------------------------------------------------------\n\ncell61_models = dict(\n    zip(\n        cell61_plan[\"target\"],\n        cell61_plan[\"selected_model\"]\n    )\n)\n\ncell63_models = dict(\n    zip(\n        cell63_plan[\"target\"],\n        cell63_plan[\"selected_model\"]\n    )\n)\n\nif cell61_models != cell63_models:\n    raise RuntimeError(\n        \"Cell-61 and Cell-63 selected-model assignments disagree.\"\n    )\n\nprint(\"Cell-61 vs Cell-63 selected-model consistency: PASS\")\n\n# ----------------------------------------------------------------------\n# Thresholds come ONLY from Cell 61\n# ----------------------------------------------------------------------\n\nthreshold_map = dict(\n    zip(\n        cell61_plan[\"target\"],\n        pd.to_numeric(\n            cell61_plan[\"selected_threshold\"],\n            errors=\"raise\"\n        )\n    )\n)\n\nif set(threshold_map.keys()) != set(EXPECTED_TARGETS):\n    raise RuntimeError(\n        \"Cell-61 threshold coverage does not match TARGETS.\"\n    )\n\nfor target in EXPECTED_TARGETS:\n\n    threshold = float(\n        threshold_map[target]\n    )\n\n    if not np.isfinite(threshold):\n        raise RuntimeError(\n            f\"Threshold for '{target}' is not finite.\"\n        )\n\n    if not 0.0 <= threshold <= 1.0:\n        raise RuntimeError(\n            f\"Threshold for '{target}' is outside [0,1]: \"\n            f\"{threshold}\"\n        )\n\nprint(\"OOF-derived thresholds loaded from Cell 61: PASS\")\nprint(\"Selected threshold range: PASS\")\n\n# ----------------------------------------------------------------------\n# Validate the exact three-fold mapping from Cell 63\n# ----------------------------------------------------------------------\n\nexpected_model_names = {\n    \"baseline\": [\n        \"baseline_fold0\",\n        \"baseline_fold1\",\n        \"baseline_fold2\",\n    ],\n    \"fine_tuned\": [\n        \"finetuned_fold0\",\n        \"finetuned_fold1\",\n        \"finetuned_fold2\",\n    ],\n}\n\nfor _, row in cell63_plan.iterrows():\n\n    selected_model = row[\"selected_model\"]\n\n    if selected_model not in expected_model_names:\n        raise RuntimeError(\n            f\"Unexpected selected_model '{selected_model}' \"\n            f\"for target '{row['target']}'.\"\n        )\n\n    expected_folds = expected_model_names[\n        selected_model\n    ]\n\n    actual_folds = [\n        row[\"fold0_model\"],\n        row[\"fold1_model\"],\n        row[\"fold2_model\"],\n    ]\n\n    if actual_folds != expected_folds:\n        raise RuntimeError(\n            f\"Invalid fold mapping for target \"\n            f\"'{row['target']}'.\\n\"\n            f\"Expected: {expected_folds}\\n\"\n            f\"Found:    {actual_folds}\"\n        )\n\nprint(\"Three-fold model mapping: PASS\")\nprint(\"Baseline/fine-tuned family mapping: PASS\")\n\n# ======================================================================\n# FINAL PROBABILITY -> BINARY SUBMISSION CONSISTENCY\n# ======================================================================\n\nfor target in EXPECTED_TARGETS:\n\n    threshold = float(\n        threshold_map[target]\n    )\n\n    expected_binary = (\n        probabilities[target]\n        .to_numpy(dtype=np.float32)\n        >= threshold\n    ).astype(int)\n\n    actual_binary = (\n        submission[target]\n        .to_numpy(dtype=int)\n    )\n\n    if not np.array_equal(\n        expected_binary,\n        actual_binary\n    ):\n        raise RuntimeError(\n            f\"Final submission mismatch for target \"\n            f\"'{target}'. \"\n            f\"Submission values do not equal the Cell-61 \"\n            f\"OOF-derived threshold applied to Cell-67 probabilities.\"\n        )\n\nprint(\"Probability-to-binary consistency: PASS\")\nprint(\"OOF-derived threshold application: PASS\")\n\n# ----------------------------------------------------------------------\n# Final submission preview\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL SUBMISSION\")\nprint(\"=\" * 70)\n\nprint(submission.to_string(index=False))\n\n# ----------------------------------------------------------------------\n# Final verdict\n# ----------------------------------------------------------------------\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 68R VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Final submission exists: PASS\")\nprint(\"Final submission schema: PASS\")\nprint(\"Target ordering: PASS\")\nprint(\"3/3 test studies covered: PASS\")\nprint(\"StudyInstanceUID uniqueness: PASS\")\nprint(\"All 12 targets present: PASS\")\nprint(\"No missing target values: PASS\")\nprint(\"Binary prediction validity: PASS\")\nprint(\"Probability file validity: PASS\")\nprint(\"Cell-61 OOF-backed model selection: PASS\")\nprint(\"Cell-61 OOF-derived thresholds: PASS\")\nprint(\"Cell-63 three-fold mapping: PASS\")\nprint(\"Probability-to-submission consistency: PASS\")\n\nprint()\nprint(\"No training performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No test labels used.\")\nprint(\"No additional inference performed.\")\n\nprint()\nprint(\"Final submission:\")\nprint(FINAL_SUBMISSION_PATH)\n\nprint()\nprint(\"Final probability file:\")\nprint(PROBABILITY_PATH)\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 68R COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T12:10:52.882128Z","iopub.execute_input":"2026-08-11T12:10:52.882479Z","iopub.status.idle":"2026-08-11T12:10:52.94176Z","shell.execute_reply.started":"2026-08-11T12:10:52.882449Z","shell.execute_reply":"2026-08-11T12:10:52.940912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ======================================================================\n# CELL 69 - FINAL KAGGLE PROBABILITY SUBMISSION\n# ======================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\n\nprint(\"=\" * 70)\nprint(\"CELL 69 - FINAL KAGGLE PROBABILITY SUBMISSION\")\nprint(\"=\" * 70)\n\n# ----------------------------------------------------------------------\n# REQUIRED CONFIGURATION\n# ----------------------------------------------------------------------\n\nTARGETS = [\n    \"ACL\",\n    \"MCL\",\n    \"Medial Meniscus\",\n    \"Lateral Meniscus\",\n    \"Medial OA\",\n    \"Lateral OA\",\n    \"PF OA\",\n    \"Effusion\",\n    \"Synovitis\",\n    \"Baker's\",\n    \"Contusion\",\n    \"Fracture\",\n]\n\nEXPECTED_COLUMNS = [\n    \"StudyInstanceUID\"\n] + TARGETS\n\nAUDIT_DIR = \"/kaggle/working/rsna_knee_audit\"\n\nPROBABILITY_FILE = os.path.join(\n    AUDIT_DIR,\n    \"cell67_final_test_probabilities.csv\"\n)\n\nFINAL_SUBMISSION = \"/kaggle/working/submission.csv\"\n\n# ----------------------------------------------------------------------\n# FILE EXISTENCE\n# ----------------------------------------------------------------------\n\nif not os.path.isfile(PROBABILITY_FILE):\n    raise FileNotFoundError(\n        \"Validated probability file not found: \"\n        + PROBABILITY_FILE\n    )\n\nprint(\"Validated probability file: EXISTS\")\n\n# ----------------------------------------------------------------------\n# LOAD CONTINUOUS PROBABILITIES\n# ----------------------------------------------------------------------\n\nprob_df = pd.read_csv(PROBABILITY_FILE)\n\nprint(\"Probability input shape:\", prob_df.shape)\n\n# ----------------------------------------------------------------------\n# COLUMN VALIDATION\n# ----------------------------------------------------------------------\n\nif list(prob_df.columns) != EXPECTED_COLUMNS:\n    raise RuntimeError(\n        \"Probability file column ordering/schema mismatch.\\n\"\n        \"Expected: \"\n        + str(EXPECTED_COLUMNS)\n        + \"\\nFound: \"\n        + str(list(prob_df.columns))\n    )\n\nprint(\"Kaggle submission schema: PASS\")\nprint(\"Column ordering: PASS\")\nprint(\"Column count:\", len(prob_df.columns))\n\n# ----------------------------------------------------------------------\n# ROW VALIDATION\n# ----------------------------------------------------------------------\n\nif len(prob_df) != 3:\n    raise RuntimeError(\n        \"Expected 3 example test studies, found \"\n        + str(len(prob_df))\n    )\n\nprint(\"Submission row count:\", len(prob_df))\nprint(\"Expected test studies: 3\")\nprint(\"Row count: PASS\")\n\n# ----------------------------------------------------------------------\n# STUDY ID VALIDATION\n# ----------------------------------------------------------------------\n\nif prob_df[\"StudyInstanceUID\"].isna().any():\n    raise RuntimeError(\n        \"StudyInstanceUID contains missing values.\"\n    )\n\nif prob_df[\"StudyInstanceUID\"].duplicated().any():\n    raise RuntimeError(\n        \"Duplicate StudyInstanceUID detected.\"\n    )\n\nprint(\"StudyInstanceUID completeness: PASS\")\nprint(\"StudyInstanceUID uniqueness: PASS\")\n\n# ----------------------------------------------------------------------\n# TARGET NUMERIC VALIDATION\n# ----------------------------------------------------------------------\n\nfor target in TARGETS:\n\n    if not pd.api.types.is_numeric_dtype(prob_df[target]):\n        raise RuntimeError(\n            f\"Target column is not numeric: {target}\"\n        )\n\nprint(\"Target numeric types: PASS\")\n\n# ----------------------------------------------------------------------\n# PROBABILITY VALIDATION\n# ----------------------------------------------------------------------\n\ntarget_values = prob_df[TARGETS].to_numpy(\n    dtype=np.float64\n)\n\nif not np.isfinite(target_values).all():\n    raise RuntimeError(\n        \"Probability file contains NaN or Inf values.\"\n    )\n\nprint(\"NaN/Inf validation: PASS\")\n\nif (target_values < 0).any():\n    raise RuntimeError(\n        \"Probability below 0 detected.\"\n    )\n\nif (target_values > 1).any():\n    raise RuntimeError(\n        \"Probability above 1 detected.\"\n    )\n\nprint(\"Probability range [0,1]: PASS\")\n\n# ----------------------------------------------------------------------\n# IMPORTANT:\n# DO NOT THRESHOLD THE PROBABILITIES.\n# ROC-AUC REQUIRES CONTINUOUS CONFIDENCE SCORES.\n# ----------------------------------------------------------------------\n\nif np.all(\n    np.isin(\n        target_values,\n        [0.0, 1.0]\n    )\n):\n    raise RuntimeError(\n        \"All target values are binary 0/1. \"\n        \"Kaggle requires continuous confidence scores. \"\n        \"Use the Cell 67 probability file, not the thresholded submission.\"\n    )\n\nprint(\"Continuous confidence scores: PASS\")\n\n# ----------------------------------------------------------------------\n# SAVE EXACT KAGGLE FILENAME\n# ----------------------------------------------------------------------\n\nprob_df.to_csv(\n    FINAL_SUBMISSION,\n    index=False\n)\n\nif not os.path.isfile(FINAL_SUBMISSION):\n    raise RuntimeError(\n        \"Final submission.csv was not created.\"\n    )\n\nprint(\"Final submission file: CREATED\")\nprint(\"Path:\", FINAL_SUBMISSION)\n\n# ----------------------------------------------------------------------\n# RELOAD AND FINAL AUDIT\n# ----------------------------------------------------------------------\n\nfinal_df = pd.read_csv(FINAL_SUBMISSION)\n\nif list(final_df.columns) != EXPECTED_COLUMNS:\n    raise RuntimeError(\n        \"Reloaded submission schema mismatch.\"\n    )\n\nif final_df.shape != (3, 13):\n    raise RuntimeError(\n        \"Unexpected final submission shape: \"\n        + str(final_df.shape)\n    )\n\nfinal_values = final_df[TARGETS].to_numpy(\n    dtype=np.float64\n)\n\nif not np.isfinite(final_values).all():\n    raise RuntimeError(\n        \"Final submission contains NaN/Inf.\"\n    )\n\nif (final_values < 0).any() or (final_values > 1).any():\n    raise RuntimeError(\n        \"Final submission contains values outside [0,1].\"\n    )\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL KAGGLE SUBMISSION AUDIT\")\nprint(\"=\" * 70)\n\nprint(\"Shape:\", final_df.shape)\nprint(\"Schema: PASS\")\nprint(\"Column ordering: PASS\")\nprint(\"StudyInstanceUID completeness: PASS\")\nprint(\"StudyInstanceUID uniqueness: PASS\")\nprint(\"12 target columns: PASS\")\nprint(\"Continuous confidence scores: PASS\")\nprint(\"Probability range [0,1]: PASS\")\nprint(\"NaN/Inf validation: PASS\")\n\nprint()\nprint(\"Final submission preview:\")\nprint(final_df.to_string(index=False))\n\nprint()\nprint(\"=\" * 70)\nprint(\"CELL 69 VERDICT\")\nprint(\"=\" * 70)\n\nprint(\"Kaggle ROC-AUC probability format: PASS\")\nprint(\"Continuous probabilities preserved: PASS\")\nprint(\"Thresholding NOT applied: PASS\")\nprint(\"Submission filename submission.csv: PASS\")\nprint(\"Final submission schema: PASS\")\nprint(\"Final submission ready for Kaggle upload: PASS\")\n\nprint()\nprint(\"Final submission:\")\nprint(FINAL_SUBMISSION)\n\nprint()\nprint(\"No training performed.\")\nprint(\"No inference performed.\")\nprint(\"No checkpoint modified.\")\nprint(\"No test labels used.\")\nprint(\"=\" * 70)\nprint(\"CELL 69 COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-11T12:20:47.380633Z","iopub.execute_input":"2026-08-11T12:20:47.381041Z","iopub.status.idle":"2026-08-11T12:20:47.418958Z","shell.execute_reply.started":"2026-08-11T12:20:47.381009Z","shell.execute_reply":"2026-08-11T12:20:47.417629Z"}},"outputs":[],"execution_count":null}]}