{"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":"markdown","source":"<div style=\"display: flex; justify-content: space-between; align-items: flex-start;\">\n    <div style=\"text-align: left;\">\n        <p style=\"color:#FFD700; font-size: 15px; font-weight: bold; margin-bottom: 1px; text-align: left;\">Published on  September 08, 2026</p>\n        <h4 style=\"color:#4B0082; font-weight: bold; text-align: left; margin-top: 6px;\">Author: Jocelyn C. Dumlao</h4>\n        <p style=\"font-size: 17px; line-height: 1.7; color: #333; text-align: center; margin-top: 20px;\"></p>\n        <a href=\"https://www.linkedin.com/in/jocelyn-dumlao-168921a8/\" target=\"_blank\" style=\"display: inline-block; background-color: #003f88; color: #fff; text-decoration: none; padding: 5px 10px; border-radius: 10px; margin: 15px;\">LinkedIn</a>\n        <a href=\"https://github.com/jcdumlao14\" target=\"_blank\" style=\"display: inline-block; background-color: transparent; color: #059c99; text-decoration: none; padding: 5px 10px; border-radius: 10px; margin: 15px; border: 2px solid #007bff;\">GitHub</a>\n        <a href=\"https://www.youtube.com/@CogniCraftedMinds\" target=\"_blank\" style=\"display: inline-block; background-color: #ff0054; color: #fff; text-decoration: none; padding: 5px 10px; border-radius: 10px; margin: 15px;\">YouTube</a>\n        <a href=\"https://www.kaggle.com/jocelyndumlao\" target=\"_blank\" style=\"display: inline-block; background-color: #3a86ff; color: #fff; text-decoration: none; padding: 5px 10px; border-radius: 10px; margin: 15px;\">Kaggle</a>\n    </div>\n</div>","metadata":{}},{"cell_type":"markdown","source":"# <div style=\"color:white;display:inline-block;border-radius:4px;background-color:#065535 ;font-family:Nexa;overflow:hidden\"><p style=\"padding:8px;color:white;overflow:hidden;font-size:85%;letter-spacing:0.5px;margin:0;border: 6px groove #e4c155;\"><b> </b>Setup</p></div>\n","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# 0. STRICT OFFLINE MODE\n# MUST RUN BEFORE importing timm / huggingface_hub\n# ============================================================\n\nimport os\n\nos.environ[\"HF_HUB_OFFLINE\"] = \"1\"\nos.environ[\"HF_DATASETS_OFFLINE\"] = \"1\"\nos.environ[\"TRANSFORMERS_OFFLINE\"] = \"1\"\n\nos.environ[\"HF_HUB_DISABLE_TELEMETRY\"] = \"1\"\nos.environ[\"HF_HUB_DISABLE_XET\"] = \"1\"\nos.environ[\"HF_HUB_ENABLE_HF_TRANSFER\"] = \"0\"\n\nprint(\"=\" * 70)\nprint(\"OFFLINE MODE ENABLED\")\nprint(\"=\" * 70)\nprint(\"HF_HUB_OFFLINE       =\", os.environ[\"HF_HUB_OFFLINE\"])\nprint(\"HF_DATASETS_OFFLINE  =\", os.environ[\"HF_DATASETS_OFFLINE\"])\nprint(\"TRANSFORMERS_OFFLINE =\", os.environ[\"TRANSFORMERS_OFFLINE\"])\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T08:01:45.939961Z","iopub.execute_input":"2026-09-14T08:01:45.940858Z","iopub.status.idle":"2026-09-14T08:01:45.958812Z","shell.execute_reply.started":"2026-09-14T08:01:45.940822Z","shell.execute_reply":"2026-09-14T08:01:45.958014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 1. OFFLINE ENVIRONMENT — MUST COME FIRST\n# ============================================================\n\nimport os\n\n# ------------------------------------------------------------\n# Hugging Face / Transformers offline mode\n# MUST be configured before importing timm.\n# ------------------------------------------------------------\n\nos.environ[\"HF_HUB_OFFLINE\"] = \"1\"\nos.environ[\"TRANSFORMERS_OFFLINE\"] = \"1\"\nos.environ[\"HF_HUB_DISABLE_TELEMETRY\"] = \"1\"\n\n# Disable Hugging Face Xet/network-related behavior.\nos.environ[\"HF_HUB_DISABLE_XET\"] = \"1\"\n\n# Avoid token/network lookup.\nos.environ[\"HF_HUB_DISABLE_IMPLICIT_TOKEN\"] = \"1\"\n\n# Force local cache locations where possible.\nos.environ.setdefault(\n    \"HF_HOME\",\n    \"/kaggle/working/huggingface\"\n)\n\nos.environ.setdefault(\n    \"TRANSFORMERS_CACHE\",\n    \"/kaggle/working/huggingface\"\n)\n\nos.environ.setdefault(\n    \"HF_DATASETS_OFFLINE\",\n    \"1\"\n)\n\nprint(\"=\" * 70)\nprint(\"STRICT OFFLINE MODE\")\nprint(\"=\" * 70)\n\nprint(\n    \"HF_HUB_OFFLINE:\",\n    os.environ.get(\"HF_HUB_OFFLINE\")\n)\n\nprint(\n    \"TRANSFORMERS_OFFLINE:\",\n    os.environ.get(\"TRANSFORMERS_OFFLINE\")\n)\n\nprint(\n    \"HF_HUB_DISABLE_XET:\",\n    os.environ.get(\"HF_HUB_DISABLE_XET\")\n)\n\nprint()\n\n\n# ============================================================\n# 2. IMPORTS\n# ============================================================\n\nimport gc\nimport math\nimport pickle\nimport time\nimport warnings\n\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n# ------------------------------------------------------------\n# timm is imported ONLY AFTER offline environment variables\n# have been configured.\n# ------------------------------------------------------------\n\nimport timm\n\nfrom pydicom.pixel_data_handlers.util import (\n    apply_modality_lut,\n)\n\nwarnings.filterwarnings(\"ignore\")\n\n\n# ============================================================\n# 3. ENVIRONMENT\n# ============================================================\n\nDEVICE = torch.device(\n    \"cuda\"\n    if torch.cuda.is_available()\n    else \"cpu\"\n)\n\nif DEVICE.type == \"cuda\":\n\n    torch.backends.cudnn.benchmark = True\n\n    try:\n\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True\n\n    except Exception:\n\n        pass\n\n\nprint(\"=\" * 70)\nprint(\"RSNA KNEE ABNORMALITY — OFFLINE V6\")\nprint(\"=\" * 70)\n\nprint(\n    f\"Device: {DEVICE}\"\n)\n\nif DEVICE.type == \"cuda\":\n\n    print(\n        f\"GPU: \"\n        f\"{torch.cuda.get_device_name(0)}\"\n    )\n\nprint()\n\n\n# ============================================================\n# 4. CONFIGURATION\n# ============================================================\n\nIMAGE_SIZE = 384\n\nCROP_MM = 140.0\n\nTOTAL_SLICES = 64\n\nDEFAULT_ARCH = (\n    \"coatnet_rmlp_2_rw_384.sw_in12k_ft_in1k\"\n)\n\nDEFAULT_RESOLUTION = 384\n\nOUTPUT_FILE = (\n    \"/kaggle/working/submission.csv\"\n)\n\nCACHE_FILE = (\n    \"/kaggle/working/\"\n    \"rsna_knee_study_lookup_cache_v6.pkl\"\n)\n\nCACHE_VERSION = (\n    \"rsna_knee_lookup_v6_strict_offline\"\n)\n\nPROGRESS_EVERY_STUDIES = 25\n\nPROGRESS_EVERY_FILES = 1000\n\n# ------------------------------------------------------------\n# Optional local checkpoint.\n#\n# If you later attach a trained RSNA checkpoint as a Kaggle\n# Dataset, put its path here.\n#\n# Example:\n#\n# MODEL_CHECKPOINT = (\n#     \"/kaggle/input/my-rsna-model/\"\n#     \"best_model.pth\"\n# )\n#\n# Leave as None for checkpoint-free mode.\n# ------------------------------------------------------------\n\nMODEL_CHECKPOINT = None\n\n\n# ============================================================\n# 5. LABELS\n# ============================================================\n\nLABELS = [\n\n    \"ACL\",\n\n    \"MCL\",\n\n    \"Medial Meniscus\",\n\n    \"Lateral Meniscus\",\n\n    \"Medial OA\",\n\n    \"Lateral OA\",\n\n    \"PF OA\",\n\n    \"Effusion\",\n\n    \"Synovitis\",\n\n    \"Baker's\",\n\n    \"Contusion\",\n\n    \"Fracture\",\n]\n\nNUM_CLASSES = len(\n    LABELS\n)\n\n\n# ============================================================\n# 6. IMAGENET NORMALIZATION\n# ============================================================\n\nIMAGENET_MEAN = np.array(\n    [\n        0.485,\n        0.456,\n        0.406,\n    ],\n    dtype=np.float32,\n)\n\nIMAGENET_STD = np.array(\n    [\n        0.229,\n        0.224,\n        0.225,\n    ],\n    dtype=np.float32,\n)\n\n\n# ============================================================\n# 7. ANATOMICAL SERIES SLOTS\n# ============================================================\n\nSERIES_SLOTS = [\n\n    (\n        \"Sagittal\",\n        1,\n        18,\n    ),\n\n    (\n        \"Sagittal\",\n        0,\n        14,\n    ),\n\n    (\n        \"Coronal\",\n        1,\n        12,\n    ),\n\n    (\n        \"Coronal\",\n        0,\n        8,\n    ),\n\n    (\n        \"Axial\",\n        -1,\n        12,\n    ),\n]\n\n\n# ============================================================\n# 8. TTA CONFIGURATIONS\n# ============================================================\n\nTTA_CONFIGS = [\n\n    {\n        \"name\": \"wide_02_98\",\n        \"lo\": 0.02,\n        \"hi\": 0.98,\n        \"k\": 64,\n        \"weight\": 0.28,\n    },\n\n    {\n        \"name\": \"wide_04_96\",\n        \"lo\": 0.04,\n        \"hi\": 0.96,\n        \"k\": 64,\n        \"weight\": 0.24,\n    },\n\n    {\n        \"name\": \"balanced_06_94\",\n        \"lo\": 0.06,\n        \"hi\": 0.94,\n        \"k\": 60,\n        \"weight\": 0.20,\n    },\n\n    {\n        \"name\": \"central_10_90\",\n        \"lo\": 0.10,\n        \"hi\": 0.90,\n        \"k\": 56,\n        \"weight\": 0.16,\n    },\n\n    {\n        \"name\": \"focused_15_85\",\n        \"lo\": 0.15,\n        \"hi\": 0.85,\n        \"k\": 48,\n        \"weight\": 0.12,\n    },\n]\n\n\n# ============================================================\n# 9. TARGET RANK WEIGHTS\n# ============================================================\n\nTARGET_RANK_WEIGHTS = {\n\n    \"ACL\": 0.70,\n\n    \"MCL\": 0.55,\n\n    \"Medial Meniscus\": 0.65,\n\n    \"Lateral Meniscus\": 0.72,\n\n    \"Medial OA\": 0.60,\n\n    \"Lateral OA\": 0.68,\n\n    \"PF OA\": 0.58,\n\n    \"Effusion\": 0.52,\n\n    \"Synovitis\": 0.58,\n\n    \"Baker's\": 0.55,\n\n    \"Contusion\": 0.60,\n\n    \"Fracture\": 0.70,\n}\n\n\n# ============================================================\n# 10. LOCATE COMPETITION ROOT\n# ============================================================\n\ndef locate_competition_root():\n\n    candidates = [\n\n        Path(\n            \"/kaggle/input/competitions/\"\n            \"rsna-knee-abnormality-detection\"\n        ),\n\n        Path(\n            \"/kaggle/input/\"\n            \"rsna-knee-abnormality-detection\"\n        ),\n\n        Path(\"data\"),\n\n        Path(\".\"),\n    ]\n\n    for root in candidates:\n\n        if not root.exists():\n\n            continue\n\n        if (\n\n            (root / \"test.csv\").is_file()\n\n            and\n\n            (\n                root /\n                \"test_series.csv\"\n            ).is_file()\n\n            and\n\n            (\n                root /\n                \"test_series\"\n            ).is_dir()\n\n        ):\n\n            print(\n                \"Using competition root:\"\n            )\n\n            print(\n                f\"  {root}\"\n            )\n\n            return root\n\n    # --------------------------------------------------------\n    # Fallback search.\n    # --------------------------------------------------------\n\n    input_root = Path(\n        \"/kaggle/input\"\n    )\n\n    if input_root.exists():\n\n        print(\n            \"Searching /kaggle/input \"\n            \"for test.csv + test_series...\"\n        )\n\n        for csv_path in input_root.glob(\n            \"**/test.csv\"\n        ):\n\n            root = (\n                csv_path.parent\n            )\n\n            if (\n\n                (\n                    root /\n                    \"test_series\"\n                ).is_dir()\n\n                and\n\n                (\n                    root /\n                    \"test_series.csv\"\n                ).is_file()\n\n            ):\n\n                print(\n                    \"Discovered competition root:\"\n                )\n\n                print(\n                    f\"  {root}\"\n                )\n\n                return root\n\n    raise FileNotFoundError(\n\n        \"\\nCould not locate the RSNA competition root.\\n\"\n\n        \"\\nExpected:\\n\"\n\n        \"  test.csv\\n\"\n\n        \"  test_series.csv\\n\"\n\n        \"  test_series/\\n\"\n    )\n\n\n# ============================================================\n# 11. LOCATE TEST SERIES\n# ============================================================\n\ndef locate_test_series(root):\n\n    candidates = [\n\n        root / \"test_series\",\n\n        root / \"test_images\",\n\n        root / \"test\",\n\n        root / \"images\" / \"test\",\n    ]\n\n    official = (\n        root /\n        \"test_series\"\n    )\n\n    if official.is_dir():\n\n        print(\n            \"Using official test DICOM directory:\"\n        )\n\n        print(\n            f\"  {official}\"\n        )\n\n        return official\n\n    for path in candidates:\n\n        if path.is_dir():\n\n            print(\n                \"Using discovered test DICOM directory:\"\n            )\n\n            print(\n                f\"  {path}\"\n            )\n\n            return path\n\n    raise FileNotFoundError(\n\n        \"\\nCould not locate test DICOM directory.\\n\"\n\n        \"\\nExpected:\\n\"\n\n        f\"  {root / 'test_series'}\\n\"\n    )\n\n\n# ============================================================\n# 12. LOAD TEST CSV\n# ============================================================\n\ndef load_test_csv(root):\n\n    test_csv = (\n        root /\n        \"test.csv\"\n    )\n\n    if not test_csv.is_file():\n\n        matches = list(\n            root.glob(\n                \"**/test.csv\"\n            )\n        )\n\n        if not matches:\n\n            raise FileNotFoundError(\n                \"test.csv not found.\"\n            )\n\n        test_csv = matches[0]\n\n    df = pd.read_csv(\n        test_csv\n    )\n\n    if df.empty:\n\n        raise ValueError(\n            \"test.csv is empty.\"\n        )\n\n    if (\n        \"StudyInstanceUID\"\n        not in df.columns\n    ):\n\n        raise ValueError(\n            \"test.csv does not contain \"\n            \"'StudyInstanceUID'.\"\n        )\n\n    print(\n        \"Test CSV:\"\n    )\n\n    print(\n        f\"  {test_csv}\"\n    )\n\n    print(\n        f\"Test studies: \"\n        f\"{len(df):,}\"\n    )\n\n    return (\n        df,\n        \"StudyInstanceUID\",\n        test_csv,\n    )\n\n\n# ============================================================\n# 13. LOAD TEST SERIES CSV\n# ============================================================\n\ndef load_test_series_csv(root):\n\n    test_series_csv = (\n        root /\n        \"test_series.csv\"\n    )\n\n    if not test_series_csv.is_file():\n\n        matches = list(\n            root.glob(\n                \"**/test_series.csv\"\n            )\n        )\n\n        if not matches:\n\n            raise FileNotFoundError(\n                \"test_series.csv not found.\"\n            )\n\n        test_series_csv = matches[0]\n\n    df = pd.read_csv(\n        test_series_csv\n    )\n\n    required = [\n\n        \"StudyInstanceUID\",\n\n        \"SeriesInstanceUID\",\n\n        \"Fluid_Sensitive\",\n\n        \"Anatomical_Plane\",\n    ]\n\n    missing = [\n\n        col\n\n        for col in required\n\n        if col not in df.columns\n    ]\n\n    if missing:\n\n        raise ValueError(\n\n            \"test_series.csv is missing \"\n\n            f\"columns: {missing}\"\n        )\n\n    print(\n        \"Test series CSV:\"\n    )\n\n    print(\n        f\"  {test_series_csv}\"\n    )\n\n    print(\n        f\"Test series: \"\n        f\"{len(df):,}\"\n    )\n\n    print()\n\n    print(\n        \"Series metadata columns:\"\n    )\n\n    print(\n        f\"  {list(df.columns)}\"\n    )\n\n    return (\n        df,\n        test_series_csv,\n    )\n\n\n# ============================================================\n# 14. DICOM HEADER TAGS\n# ============================================================\n\nDICOM_TAGS = [\n\n    \"StudyInstanceUID\",\n\n    \"SeriesInstanceUID\",\n\n    \"SOPInstanceUID\",\n\n    \"SeriesDescription\",\n\n    \"ProtocolName\",\n\n    \"SequenceName\",\n\n    \"ImageOrientationPatient\",\n\n    \"ImagePositionPatient\",\n\n    \"PixelSpacing\",\n\n    \"InstanceNumber\",\n\n    \"SliceThickness\",\n\n    \"SpacingBetweenSlices\",\n\n    \"PatientPosition\",\n\n    \"Modality\",\n\n    \"BodyPartExamined\",\n]\n\n\n# ============================================================\n# 15. SAFE HELPERS\n# ============================================================\n\ndef safe_float_array(value):\n\n    if value is None:\n\n        return None\n\n    try:\n\n        return np.asarray(\n\n            [\n                float(x)\n                for x in value\n            ],\n\n            dtype=np.float32,\n        )\n\n    except Exception:\n\n        return None\n\n\ndef safe_float(\n    value,\n    default=None,\n):\n\n    try:\n\n        return float(value)\n\n    except Exception:\n\n        return default\n\n\ndef safe_int(\n    value,\n    default=None,\n):\n\n    try:\n\n        return int(value)\n\n    except Exception:\n\n        return default\n\n\ndef get_string(\n    ds,\n    name,\n):\n\n    try:\n\n        value = getattr(\n            ds,\n            name,\n            None,\n        )\n\n        if value is None:\n\n            return \"\"\n\n        return str(value)\n\n    except Exception:\n\n        return \"\"\n\n\n# ============================================================\n# 16. ESTIMATE PLANE\n# ============================================================\n\ndef estimate_plane(ds):\n\n    try:\n\n        iop = safe_float_array(\n\n            getattr(\n\n                ds,\n\n                \"ImageOrientationPatient\",\n\n                None,\n            )\n        )\n\n        if (\n\n            iop is None\n\n            or\n\n            len(iop) != 6\n\n        ):\n\n            return \"\"\n\n        row = iop[:3]\n\n        col = iop[3:]\n\n        normal = np.cross(\n            row,\n            col,\n        )\n\n        axis = int(\n            np.argmax(\n                np.abs(normal)\n            )\n        )\n\n        if axis == 0:\n\n            return \"Sagittal\"\n\n        if axis == 1:\n\n            return \"Coronal\"\n\n        if axis == 2:\n\n            return \"Axial\"\n\n    except Exception:\n\n        pass\n\n    return \"\"\n\n\n# ============================================================\n# 17. READ LIGHT DICOM METADATA\n# ============================================================\n\ndef read_dicom_metadata(path):\n\n    try:\n\n        ds = pydicom.dcmread(\n\n            str(path),\n\n            stop_before_pixels=True,\n\n            specific_tags=DICOM_TAGS,\n\n            force=True,\n        )\n\n        iop = safe_float_array(\n\n            getattr(\n\n                ds,\n\n                \"ImageOrientationPatient\",\n\n                None,\n            )\n        )\n\n        ipp = safe_float_array(\n\n            getattr(\n\n                ds,\n\n                \"ImagePositionPatient\",\n\n                None,\n            )\n        )\n\n        spacing = safe_float_array(\n\n            getattr(\n\n                ds,\n\n                \"PixelSpacing\",\n\n                None,\n            )\n        )\n\n        if (\n\n            spacing is None\n\n            or\n\n            len(spacing) < 2\n\n        ):\n\n            spacing = np.array(\n\n                [\n                    0.5,\n                    0.5,\n                ],\n\n                dtype=np.float32,\n            )\n\n        return {\n\n            \"path\":\n                str(path),\n\n            \"study_uid\":\n                get_string(\n                    ds,\n                    \"StudyInstanceUID\",\n                ),\n\n            \"series_uid\":\n                get_string(\n                    ds,\n                    \"SeriesInstanceUID\",\n                ),\n\n            \"sop_uid\":\n                get_string(\n                    ds,\n                    \"SOPInstanceUID\",\n                ),\n\n            \"series_description\":\n                get_string(\n                    ds,\n                    \"SeriesDescription\",\n                ),\n\n            \"protocol_name\":\n                get_string(\n                    ds,\n                    \"ProtocolName\",\n                ),\n\n            \"sequence_name\":\n                get_string(\n                    ds,\n                    \"SequenceName\",\n                ),\n\n            \"plane\":\n                estimate_plane(ds),\n\n            \"image_orientation\":\n                (\n                    iop.tolist()\n                    if iop is not None\n                    else None\n                ),\n\n            \"image_position\":\n                (\n                    ipp.tolist()\n                    if ipp is not None\n                    else None\n                ),\n\n            \"pixel_spacing\":\n                spacing[:2].tolist(),\n\n            \"instance_number\":\n                safe_int(\n                    getattr(\n                        ds,\n                        \"InstanceNumber\",\n                        None,\n                    )\n                ),\n\n            \"slice_thickness\":\n                safe_float(\n                    getattr(\n                        ds,\n                        \"SliceThickness\",\n                        None,\n                    )\n                ),\n\n            \"spacing_between_slices\":\n                safe_float(\n                    getattr(\n                        ds,\n                        \"SpacingBetweenSlices\",\n                        None,\n                    )\n                ),\n\n            \"patient_position\":\n                get_string(\n                    ds,\n                    \"PatientPosition\",\n                ),\n\n            \"modality\":\n                get_string(\n                    ds,\n                    \"Modality\",\n                ),\n\n            \"body_part\":\n                get_string(\n                    ds,\n                    \"BodyPartExamined\",\n                ),\n        }\n\n    except Exception as exc:\n\n        return {\n\n            \"error\":\n                str(exc),\n\n            \"path\":\n                str(path),\n        }\n\n\n# ============================================================\n# 18. SORT SERIES\n# ============================================================\n\ndef sort_series_records(records):\n\n    if not records:\n\n        return records\n\n    orientation = None\n\n    for record in records:\n\n        value = record.get(\n            \"image_orientation\"\n        )\n\n        if value is None:\n\n            continue\n\n        try:\n\n            arr = np.asarray(\n\n                value,\n\n                dtype=np.float32,\n            )\n\n            if len(arr) == 6:\n\n                orientation = arr\n\n                break\n\n        except Exception:\n\n            pass\n\n    if orientation is not None:\n\n        try:\n\n            row = orientation[:3]\n\n            col = orientation[3:]\n\n            normal = np.cross(\n                row,\n                col,\n            )\n\n            norm = np.linalg.norm(\n                normal\n            )\n\n            if norm > 0:\n\n                normal = (\n                    normal /\n                    norm\n                )\n\n                def projection(record):\n\n                    pos = record.get(\n                        \"image_position\"\n                    )\n\n                    if pos is None:\n\n                        return float(\n                            \"inf\"\n                        )\n\n                    try:\n\n                        return float(\n\n                            np.dot(\n\n                                np.asarray(\n                                    pos,\n                                    dtype=np.float32,\n                                ),\n\n                                normal,\n                            )\n                        )\n\n                    except Exception:\n\n                        return float(\n                            \"inf\"\n                        )\n\n                return sorted(\n\n                    records,\n\n                    key=projection,\n                )\n\n        except Exception:\n\n            pass\n\n    return sorted(\n\n        records,\n\n        key=lambda x: (\n\n            x.get(\n                \"instance_number\"\n            )\n\n            if x.get(\n                \"instance_number\"\n            ) is not None\n\n            else 10**9\n        )\n    )\n\n\n# ============================================================\n# 19. CACHE SIGNATURE\n# ============================================================\n\ndef make_cache_signature(\n    test_csv,\n    test_series_csv,\n    test_studies,\n):\n\n    signature = {\n\n        \"version\":\n            CACHE_VERSION,\n\n        \"study_count\":\n            len(test_studies),\n\n        \"study_ids\":\n            tuple(\n                str(x)\n                for x in test_studies\n            ),\n    }\n\n    try:\n\n        stat = test_csv.stat()\n\n        signature[\n            \"test_csv_size\"\n        ] = stat.st_size\n\n        signature[\n            \"test_csv_mtime\"\n        ] = stat.st_mtime_ns\n\n    except Exception:\n\n        pass\n\n    try:\n\n        stat = test_series_csv.stat()\n\n        signature[\n            \"test_series_csv_size\"\n        ] = stat.st_size\n\n        signature[\n            \"test_series_csv_mtime\"\n        ] = stat.st_mtime_ns\n\n    except Exception:\n\n        pass\n\n    return signature\n\n\n# ============================================================\n# 20. LOAD CACHE\n# ============================================================\n\ndef load_lookup_cache(\n    cache_file,\n    signature,\n):\n\n    path = Path(\n        cache_file\n    )\n\n    if not path.exists():\n\n        return None\n\n    try:\n\n        print()\n        print(\n            \"=\" * 70\n        )\n\n        print(\n            \"CHECKING STUDY LOOKUP CACHE\"\n        )\n\n        print(\n            \"=\" * 70\n        )\n\n        with open(\n            path,\n            \"rb\",\n        ) as f:\n\n            payload = pickle.load(\n                f\n            )\n\n        if not isinstance(\n            payload,\n            dict,\n        ):\n\n            return None\n\n        if (\n            payload.get(\n                \"signature\"\n            )\n            != signature\n        ):\n\n            print(\n                \"Cache signature changed.\"\n            )\n\n            print(\n                \"Rebuilding lookup...\"\n            )\n\n            return None\n\n        lookup = payload.get(\n            \"lookup\"\n        )\n\n        if lookup is None:\n\n            return None\n\n        print(\n            \"Cache found:\"\n        )\n\n        print(\n            f\"  {path}\"\n        )\n\n        print(\n            f\"Cached studies: \"\n            f\"{len(lookup):,}\"\n        )\n\n        return lookup\n\n    except Exception as exc:\n\n        print(\n            \"Cache could not be loaded:\"\n        )\n\n        print(\n            f\"  {exc}\"\n        )\n\n        return None\n\n\n# ============================================================\n# 21. SAVE CACHE\n# ============================================================\n\ndef save_lookup_cache(\n    cache_file,\n    signature,\n    lookup,\n):\n\n    path = Path(\n        cache_file\n    )\n\n    try:\n\n        payload = {\n\n            \"signature\":\n                signature,\n\n            \"lookup\":\n                lookup,\n        }\n\n        temp = path.with_suffix(\n            \".tmp\"\n        )\n\n        with open(\n            temp,\n            \"wb\",\n        ) as f:\n\n            pickle.dump(\n\n                payload,\n\n                f,\n\n                protocol=\n                    pickle.HIGHEST_PROTOCOL,\n            )\n\n        os.replace(\n            temp,\n            path,\n        )\n\n        print()\n        print(\n            \"Lookup cache saved:\"\n        )\n\n        print(\n            f\"  {path}\"\n        )\n\n    except Exception as exc:\n\n        print(\n            \"Warning: could not save cache:\"\n        )\n\n        print(\n            f\"  {exc}\"\n        )\n\n\n# ============================================================\n# 22. DICOM FILE DISCOVERY\n# ============================================================\n\ndef discover_dicom_files(\n    series_dir\n):\n\n    if not series_dir.is_dir():\n\n        return []\n\n    files = []\n\n    try:\n\n        for path in series_dir.iterdir():\n\n            if not path.is_file():\n\n                continue\n\n            if (\n                path.suffix.lower()\n                == \".dcm\"\n            ):\n\n                files.append(\n                    path\n                )\n\n                continue\n\n            if path.suffix == \"\":\n\n                files.append(\n                    path\n                )\n\n    except Exception:\n\n        return []\n\n    return sorted(\n\n        files,\n\n        key=lambda p: p.name,\n    )\n\n\n# ============================================================\n# 23. CREATE SERIES LOOKUP\n# ============================================================\n\ndef create_series_lookup_optimized(\n    test_series_root,\n    test_studies,\n    test_series_df,\n    test_csv,\n    test_series_csv,\n    cache_file=CACHE_FILE,\n):\n\n    test_studies = [\n\n        str(x)\n\n        for x in test_studies\n    ]\n\n    signature = make_cache_signature(\n\n        test_csv,\n\n        test_series_csv,\n\n        test_studies,\n    )\n\n    cached = load_lookup_cache(\n\n        cache_file,\n\n        signature,\n    )\n\n    if cached is not None:\n\n        return cached\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"BUILDING TEST SERIES LOOKUP\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        f\"Studies: \"\n        f\"{len(test_studies):,}\"\n    )\n\n    print(\n        \"Test DICOM root:\"\n    )\n\n    print(\n        f\"  {test_series_root}\"\n    )\n\n    print()\n\n    lookup = {}\n\n    total_files = 0\n\n    valid_files = 0\n\n    bad_files = 0\n\n    missing_studies = 0\n\n    missing_series = 0\n\n    start_time = time.time()\n\n    metadata = (\n        test_series_df.copy()\n    )\n\n    metadata[\n        \"StudyInstanceUID\"\n    ] = metadata[\n        \"StudyInstanceUID\"\n    ].astype(str)\n\n    metadata[\n        \"SeriesInstanceUID\"\n    ] = metadata[\n        \"SeriesInstanceUID\"\n    ].astype(str)\n\n    metadata[\n        \"Fluid_Sensitive\"\n    ] = pd.to_numeric(\n\n        metadata[\n            \"Fluid_Sensitive\"\n        ],\n\n        errors=\"coerce\",\n\n    ).fillna(\n        0\n    ).astype(\n        int\n    )\n\n    metadata[\n        \"Anatomical_Plane\"\n    ] = metadata[\n        \"Anatomical_Plane\"\n    ].astype(str)\n\n    for (\n        study_index,\n        study_id,\n    ) in enumerate(\n\n        test_studies,\n\n        start=1,\n    ):\n\n        study_dir = (\n\n            test_series_root /\n            study_id\n        )\n\n        if not study_dir.is_dir():\n\n            missing_studies += 1\n\n            lookup[\n                study_id\n            ] = {}\n\n            print(\n\n                f\"[{study_index}/\"\n                f\"{len(test_studies)}] \"\n                f\"WARNING: missing study \"\n                f\"{study_id}\"\n            )\n\n            continue\n\n        study_meta = metadata[\n\n            metadata[\n                \"StudyInstanceUID\"\n            ]\n            ==\n            study_id\n        ]\n\n        series_lookup = {}\n\n        for _, row in (\n            study_meta.iterrows()\n        ):\n\n            series_uid = str(\n\n                row[\n                    \"SeriesInstanceUID\"\n                ]\n            )\n\n            series_dir = (\n\n                study_dir /\n                series_uid\n            )\n\n            if not series_dir.is_dir():\n\n                flat_candidate = (\n\n                    test_series_root /\n                    series_uid\n                )\n\n                if flat_candidate.is_dir():\n\n                    series_dir = (\n                        flat_candidate\n                    )\n\n                else:\n\n                    missing_series += 1\n\n                    continue\n\n            dcm_files = (\n                discover_dicom_files(\n                    series_dir\n                )\n            )\n\n            if not dcm_files:\n\n                missing_series += 1\n\n                continue\n\n            records = []\n\n            for dcm_path in dcm_files:\n\n                total_files += 1\n\n                record = (\n                    read_dicom_metadata(\n                        dcm_path\n                    )\n                )\n\n                if \"error\" in record:\n\n                    bad_files += 1\n\n                    continue\n\n                valid_files += 1\n\n                records.append(\n                    record\n                )\n\n                if (\n\n                    total_files\n                    %\n                    PROGRESS_EVERY_FILES\n                    ==\n                    0\n\n                ):\n\n                    elapsed = (\n                        time.time()\n                        -\n                        start_time\n                    )\n\n                    rate = (\n\n                        total_files\n                        /\n                        max(\n                            elapsed,\n                            1e-6,\n                        )\n                    )\n\n                    print(\n\n                        f\"  Files: \"\n                        f\"{total_files:,} | \"\n                        f\"Valid: \"\n                        f\"{valid_files:,} | \"\n                        f\"Bad: \"\n                        f\"{bad_files:,} | \"\n                        f\"{rate:.1f} files/s\"\n                    )\n\n            if not records:\n\n                continue\n\n            records = (\n                sort_series_records(\n                    records\n                )\n            )\n\n            plane = str(\n\n                row[\n                    \"Anatomical_Plane\"\n                ]\n            ).strip()\n\n            fluid_sensitive = int(\n\n                row[\n                    \"Fluid_Sensitive\"\n                ]\n            )\n\n            fat_suppression = int(\n\n                row.get(\n                    \"Fat_Suppression\",\n                    0,\n                )\n            )\n\n            first = records[0]\n\n            series_info = {\n\n                \"series_uid\":\n                    series_uid,\n\n                \"series_dir\":\n                    str(\n                        series_dir\n                    ),\n\n                \"plane\":\n                    plane,\n\n                \"fluid_sensitive\":\n                    fluid_sensitive,\n\n                \"fat_suppression\":\n                    fat_suppression,\n\n                \"series_description\":\n                    first.get(\n                        \"series_description\",\n                        \"\",\n                    ),\n\n                \"protocol_name\":\n                    first.get(\n                        \"protocol_name\",\n                        \"\",\n                    ),\n\n                \"sequence_name\":\n                    first.get(\n                        \"sequence_name\",\n                        \"\",\n                    ),\n\n                \"pixel_spacing\":\n                    first.get(\n                        \"pixel_spacing\",\n                        [0.5, 0.5],\n                    ),\n\n                \"body_part\":\n                    first.get(\n                        \"body_part\",\n                        \"\",\n                    ),\n\n                \"modality\":\n                    first.get(\n                        \"modality\",\n                        \"\",\n                    ),\n\n                \"records\":\n                    records,\n\n                \"image_count\":\n                    len(records),\n            }\n\n            series_lookup[\n                series_uid\n            ] = series_info\n\n        lookup[\n            study_id\n        ] = series_lookup\n\n        if (\n\n            study_index == 1\n\n            or\n\n            study_index\n            %\n            PROGRESS_EVERY_STUDIES\n            ==\n            0\n\n            or\n\n            study_index\n            ==\n            len(test_studies)\n\n        ):\n\n            elapsed = (\n                time.time()\n                -\n                start_time\n            )\n\n            rate = (\n\n                study_index\n                /\n                max(\n                    elapsed,\n                    1e-6,\n                )\n            )\n\n            remaining = (\n\n                len(test_studies)\n                -\n                study_index\n            )\n\n            eta = (\n\n                remaining\n                /\n                max(\n                    rate,\n                    1e-6,\n                )\n            )\n\n            series_count = sum(\n\n                len(x)\n\n                for x in lookup.values()\n            )\n\n            print()\n\n            print(\n\n                f\"[{study_index:,}/\"\n                f\"{len(test_studies):,}] \"\n                f\"studies | \"\n                f\"Series: \"\n                f\"{series_count:,} | \"\n                f\"Files: \"\n                f\"{total_files:,} | \"\n                f\"Bad: \"\n                f\"{bad_files:,} | \"\n                f\"ETA: \"\n                f\"{eta / 60:.2f} min\"\n            )\n\n    save_lookup_cache(\n\n        cache_file,\n\n        signature,\n\n        lookup,\n    )\n\n    elapsed = (\n        time.time()\n        -\n        start_time\n    )\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"LOOKUP COMPLETE\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        f\"Studies indexed: \"\n        f\"{len(lookup):,}\"\n    )\n\n    print(\n        f\"Missing studies: \"\n        f\"{missing_studies:,}\"\n    )\n\n    print(\n        f\"Missing series: \"\n        f\"{missing_series:,}\"\n    )\n\n    print(\n        f\"Valid DICOMs: \"\n        f\"{valid_files:,}\"\n    )\n\n    print(\n        f\"Malformed DICOMs: \"\n        f\"{bad_files:,}\"\n    )\n\n    print(\n        f\"Total files: \"\n        f\"{total_files:,}\"\n    )\n\n    print(\n        f\"Elapsed: \"\n        f\"{elapsed / 60:.2f} minutes\"\n    )\n\n    if elapsed > 0:\n\n        print(\n\n            f\"File rate: \"\n            f\"{total_files / elapsed:.1f} files/s\"\n        )\n\n    return lookup\n\n\n# ============================================================\n# 24. CHOOSE SERIES\n# ============================================================\n\ndef choose_series(\n    series_lookup,\n    plane,\n    fluid_preference=None,\n):\n\n    if not series_lookup:\n\n        return None\n\n    candidates = []\n\n    for (\n        series_uid,\n        series,\n    ) in series_lookup.items():\n\n        if (\n\n            str(\n                series.get(\n                    \"plane\",\n                    \"\",\n                )\n            ).lower()\n\n            !=\n\n            str(\n                plane\n            ).lower()\n\n        ):\n\n            continue\n\n        candidates.append(\n            series\n        )\n\n    if not candidates:\n\n        return None\n\n    if fluid_preference is not None:\n\n        desired = (\n\n            1\n\n            if fluid_preference\n\n            else 0\n        )\n\n        fluid_candidates = [\n\n            series\n\n            for series in candidates\n\n            if int(\n                series.get(\n                    \"fluid_sensitive\",\n                    0,\n                )\n            )\n            ==\n            desired\n        ]\n\n        if fluid_candidates:\n\n            candidates = (\n                fluid_candidates\n            )\n\n    return max(\n\n        candidates,\n\n        key=lambda x: int(\n\n            x.get(\n                \"image_count\",\n                0,\n            )\n        ),\n    )\n\n\n# ============================================================\n# 25. READ DICOM PIXELS\n# ============================================================\n\ndef read_dicom_pixels(path):\n\n    ds = pydicom.dcmread(\n\n        str(path),\n\n        force=True,\n    )\n\n    image = ds.pixel_array\n\n    try:\n\n        image = apply_modality_lut(\n\n            image,\n\n            ds,\n        )\n\n    except Exception:\n\n        pass\n\n    image = np.asarray(\n\n        image,\n\n        dtype=np.float32,\n    )\n\n    if (\n\n        str(\n\n            getattr(\n\n                ds,\n\n                \"PhotometricInterpretation\",\n\n                \"\",\n            )\n        ).upper()\n\n        ==\n\n        \"MONOCHROME1\"\n\n    ):\n\n        image = (\n\n            np.max(image)\n            -\n            image\n        )\n\n    image = np.nan_to_num(\n\n        image,\n\n        nan=0.0,\n\n        posinf=0.0,\n\n        neginf=0.0,\n    )\n\n    return image\n\n\n# ============================================================\n# 26. PERCENTILE NORMALIZATION\n# ============================================================\n\ndef normalize_percentile(\n    image\n):\n\n    image = np.asarray(\n\n        image,\n\n        dtype=np.float32,\n    )\n\n    valid = np.isfinite(\n        image\n    )\n\n    if not valid.any():\n\n        return np.zeros_like(\n\n            image,\n\n            dtype=np.float32,\n        )\n\n    values = image[\n        valid\n    ]\n\n    lo, hi = np.percentile(\n\n        values,\n\n        [\n            2,\n            98,\n        ],\n    )\n\n    if hi <= lo:\n\n        lo = float(\n            values.min()\n        )\n\n        hi = float(\n            values.max()\n        )\n\n    if hi <= lo:\n\n        return np.zeros_like(\n\n            image,\n\n            dtype=np.float32,\n        )\n\n    image = (\n\n        image - lo\n\n    ) / (\n\n        hi - lo\n    )\n\n    return np.clip(\n\n        image,\n\n        0.0,\n\n        1.0,\n\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 27. CENTRAL PHYSICAL CROP\n# ============================================================\n\ndef central_physical_crop(\n    image,\n    pixel_spacing,\n    crop_mm,\n):\n\n    h, w = image.shape[:2]\n\n    try:\n\n        sy = float(\n            pixel_spacing[0]\n        )\n\n        sx = float(\n            pixel_spacing[1]\n        )\n\n    except Exception:\n\n        sy = 0.5\n\n        sx = 0.5\n\n    if sy <= 0:\n\n        sy = 0.5\n\n    if sx <= 0:\n\n        sx = 0.5\n\n    crop_h = max(\n\n        1,\n\n        int(\n\n            round(\n                crop_mm / sy\n            )\n        ),\n    )\n\n    crop_w = max(\n\n        1,\n\n        int(\n\n            round(\n                crop_mm / sx\n            )\n        ),\n    )\n\n    crop_h = min(\n        crop_h,\n        h,\n    )\n\n    crop_w = min(\n        crop_w,\n        w,\n    )\n\n    y0 = max(\n\n        0,\n\n        (h - crop_h) // 2,\n    )\n\n    x0 = max(\n\n        0,\n\n        (w - crop_w) // 2,\n    )\n\n    return image[\n\n        y0:\n        y0 + crop_h,\n\n        x0:\n        x0 + crop_w,\n    ]\n\n\n# ============================================================\n# 28. RESIZE\n# ============================================================\n\ndef resize_image(\n    image,\n    size,\n):\n\n    if (\n\n        image.shape[0]\n        ==\n        size\n\n        and\n\n        image.shape[1]\n        ==\n        size\n\n    ):\n\n        return image.astype(\n            np.float32\n        )\n\n    return cv2.resize(\n\n        image,\n\n        (\n            size,\n            size,\n        ),\n\n        interpolation=\n            cv2.INTER_AREA,\n\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 29. BUILD VOLUME\n# ============================================================\n\ndef build_volume(\n    series_lookup,\n    total_slices=TOTAL_SLICES,\n):\n\n    volume = np.zeros(\n\n        (\n\n            total_slices,\n\n            IMAGE_SIZE,\n\n            IMAGE_SIZE,\n        ),\n\n        dtype=np.float32,\n    )\n\n    valid = np.zeros(\n\n        total_slices,\n\n        dtype=bool,\n    )\n\n    cursor = 0\n\n    for (\n\n        plane,\n        fluid_flag,\n        count,\n\n    ) in SERIES_SLOTS:\n\n        count = min(\n\n            int(count),\n\n            total_slices - cursor,\n        )\n\n        if count <= 0:\n\n            break\n\n        fluid_preference = (\n\n            None\n\n            if fluid_flag == -1\n\n            else bool(fluid_flag)\n        )\n\n        series = choose_series(\n\n            series_lookup,\n\n            plane,\n\n            fluid_preference,\n        )\n\n        if series is None:\n\n            cursor += count\n\n            continue\n\n        records = series.get(\n\n            \"records\",\n\n            [],\n        )\n\n        if not records:\n\n            cursor += count\n\n            continue\n\n        n = len(records)\n\n        if n == 1:\n\n            indices = np.zeros(\n\n                count,\n\n                dtype=int,\n            )\n\n        else:\n\n            indices = (\n\n                np.linspace(\n\n                    0,\n\n                    n - 1,\n\n                    count,\n                )\n\n                .round()\n\n                .astype(int)\n            )\n\n        for (\n\n            local_idx,\n            src_idx,\n\n        ) in enumerate(indices):\n\n            dst_idx = (\n\n                cursor\n                +\n                local_idx\n            )\n\n            if dst_idx >= total_slices:\n\n                break\n\n            record = records[\n                int(src_idx)\n            ]\n\n            try:\n\n                image = (\n                    read_dicom_pixels(\n                        record[\"path\"]\n                    )\n                )\n\n                image = (\n                    normalize_percentile(\n                        image\n                    )\n                )\n\n                spacing = (\n\n                    record.get(\n\n                        \"pixel_spacing\",\n\n                        [\n                            0.5,\n                            0.5,\n                        ],\n                    )\n                )\n\n                image = (\n                    central_physical_crop(\n\n                        image,\n\n                        spacing,\n\n                        CROP_MM,\n                    )\n                )\n\n                image = (\n                    resize_image(\n\n                        image,\n\n                        IMAGE_SIZE,\n                    )\n                )\n\n                volume[\n                    dst_idx\n                ] = image\n\n                valid[\n                    dst_idx\n                ] = True\n\n            except Exception:\n\n                pass\n\n        cursor += count\n\n    return (\n        volume,\n        valid,\n    )\n\n\n# ============================================================\n# 30. MAKE SLICE CENTERS\n# ============================================================\n\ndef make_centers(\n    valid_mask,\n    k,\n    lo,\n    hi,\n):\n\n    valid_indices = np.where(\n        valid_mask\n    )[0]\n\n    if len(valid_indices) == 0:\n\n        return (\n\n            np.linspace(\n\n                0,\n\n                TOTAL_SLICES - 1,\n\n                k,\n            )\n\n            .round()\n\n            .astype(int)\n        )\n\n    first = int(\n        valid_indices[0]\n    )\n\n    last = int(\n        valid_indices[-1]\n    )\n\n    lo_idx = int(\n\n        round(\n\n            first\n\n            +\n\n            lo\n            *\n            (\n                last - first\n            )\n        )\n    )\n\n    hi_idx = int(\n\n        round(\n\n            first\n\n            +\n\n            hi\n            *\n            (\n                last - first\n            )\n        )\n    )\n\n    lo_idx = max(\n\n        first,\n\n        min(\n            lo_idx,\n            last,\n        ),\n    )\n\n    hi_idx = max(\n\n        lo_idx,\n\n        min(\n            hi_idx,\n            last,\n        ),\n    )\n\n    centers = (\n\n        np.linspace(\n\n            lo_idx,\n\n            hi_idx,\n\n            k,\n        )\n\n        .round()\n\n        .astype(int)\n    )\n\n    return centers\n\n\n# ============================================================\n# 31. CREATE 2.5D RGB WINDOWS\n# ============================================================\n\ndef create_windows(\n    volume,\n    centers,\n):\n\n    windows = []\n\n    for center in centers:\n\n        center = int(\n            center\n        )\n\n        prev_idx = max(\n\n            0,\n\n            center - 1,\n        )\n\n        next_idx = min(\n\n            volume.shape[0] - 1,\n\n            center + 1,\n        )\n\n        rgb = np.stack(\n\n            [\n\n                volume[\n                    prev_idx\n                ],\n\n                volume[\n                    center\n                ],\n\n                volume[\n                    next_idx\n                ],\n            ],\n\n            axis=0,\n        )\n\n        rgb = np.clip(\n\n            rgb,\n\n            0.0,\n\n            1.0,\n\n        ).astype(\n            np.float32\n        )\n\n        rgb = (\n\n            rgb\n\n            -\n\n            IMAGENET_MEAN[\n                :,\n                None,\n                None,\n            ]\n\n        ) / (\n\n            IMAGENET_STD[\n                :,\n                None,\n                None,\n            ]\n        )\n\n        windows.append(\n            rgb\n        )\n\n    if not windows:\n\n        return np.empty(\n\n            (\n\n                0,\n\n                3,\n\n                IMAGE_SIZE,\n\n                IMAGE_SIZE,\n            ),\n\n            dtype=np.float32,\n        )\n\n    return np.stack(\n\n        windows,\n\n        axis=0,\n    )\n\n\n# ============================================================\n# 32. MODEL\n# ============================================================\n\nclass CheckpointFreeKneeModel(\n    nn.Module\n):\n\n    def __init__(\n        self,\n        arch=DEFAULT_ARCH,\n        num_classes=NUM_CLASSES,\n        pretrained=False,\n    ):\n\n        super().__init__()\n\n        print()\n        print(\n            \"Creating CoAtNet:\"\n        )\n\n        print(\n            f\"  Architecture: {arch}\"\n        )\n\n        print(\n            f\"  Pretrained: {pretrained}\"\n        )\n\n        # ----------------------------------------------------\n        # IMPORTANT:\n        #\n        # This call is offline-safe when pretrained=False.\n        #\n        # We NEVER let timm attempt an online download.\n        # ----------------------------------------------------\n\n        self.backbone = (\n            timm.create_model(\n\n                arch,\n\n                pretrained=pretrained,\n\n                num_classes=num_classes,\n\n                in_chans=3,\n            )\n        )\n\n        self.num_classes = (\n            num_classes\n        )\n\n    def forward(\n        self,\n        x,\n    ):\n\n        return self.backbone(x)\n\n\n# ============================================================\n# 33. FIND LOCAL CHECKPOINTS\n# ============================================================\n\ndef find_local_model_files():\n\n    candidates = []\n\n    # --------------------------------------------------------\n    # Explicit checkpoint first.\n    # --------------------------------------------------------\n\n    if MODEL_CHECKPOINT is not None:\n\n        path = Path(\n            MODEL_CHECKPOINT\n        )\n\n        if path.is_file():\n\n            candidates.append(\n                path\n            )\n\n    # --------------------------------------------------------\n    # Search Kaggle input for common model files.\n    # --------------------------------------------------------\n\n    kaggle_input = Path(\n        \"/kaggle/input\"\n    )\n\n    if kaggle_input.exists():\n\n        patterns = [\n\n            \"*.pth\",\n\n            \"*.pt\",\n\n            \"*.bin\",\n\n            \"*.safetensors\",\n        ]\n\n        for pattern in patterns:\n\n            try:\n\n                candidates.extend(\n\n                    kaggle_input.glob(\n                        f\"**/{pattern}\"\n                    )\n                )\n\n            except Exception:\n\n                pass\n\n    # --------------------------------------------------------\n    # Search working directory.\n    # --------------------------------------------------------\n\n    working = Path(\n        \"/kaggle/working\"\n    )\n\n    if working.exists():\n\n        patterns = [\n\n            \"*.pth\",\n\n            \"*.pt\",\n\n            \"*.bin\",\n\n            \"*.safetensors\",\n        ]\n\n        for pattern in patterns:\n\n            try:\n\n                candidates.extend(\n\n                    working.glob(\n                        f\"**/{pattern}\"\n                    )\n                )\n\n            except Exception:\n\n                pass\n\n    # --------------------------------------------------------\n    # Remove duplicates.\n    # --------------------------------------------------------\n\n    unique = []\n\n    seen = set()\n\n    for path in candidates:\n\n        try:\n\n            key = str(\n                path.resolve()\n            )\n\n        except Exception:\n\n            key = str(path)\n\n        if key in seen:\n\n            continue\n\n        seen.add(key)\n\n        unique.append(\n            path\n        )\n\n    return unique\n\n\n# ============================================================\n# 34. INSPECT CHECKPOINT\n# ============================================================\n\ndef checkpoint_looks_compatible(\n    checkpoint_path\n):\n\n    try:\n\n        checkpoint = torch.load(\n\n            checkpoint_path,\n\n            map_location=\"cpu\",\n\n            weights_only=False,\n        )\n\n    except Exception as exc:\n\n        print()\n        print(\n            \"Could not inspect checkpoint:\"\n        )\n\n        print(\n            f\"  {checkpoint_path}\"\n        )\n\n        print(\n            f\"  {exc}\"\n        )\n\n        return False\n\n    state_dict = None\n\n    if isinstance(\n        checkpoint,\n        dict,\n    ):\n\n        # Common formats.\n        for key in [\n\n            \"state_dict\",\n\n            \"model_state_dict\",\n\n            \"model\",\n\n            \"weights\",\n        ]:\n\n            if key in checkpoint:\n\n                candidate = (\n                    checkpoint[key]\n                )\n\n                if isinstance(\n                    candidate,\n                    dict,\n                ):\n\n                    state_dict = (\n                        candidate\n                    )\n\n                    break\n\n        if state_dict is None:\n\n            # Could itself be a state_dict.\n            if all(\n\n                isinstance(\n                    k,\n                    str,\n                )\n\n                for k in checkpoint.keys()\n\n            ):\n\n                state_dict = checkpoint\n\n    if state_dict is None:\n\n        print(\n            \"Checkpoint format not recognized.\"\n        )\n\n        return False\n\n    # --------------------------------------------------------\n    # Look for a classifier tensor.\n    # --------------------------------------------------------\n\n    possible_heads = []\n\n    for key, value in (\n        state_dict.items()\n    ):\n\n        if not torch.is_tensor(\n            value\n        ):\n\n            continue\n\n        if value.ndim != 2:\n\n            continue\n\n        if (\n            value.shape[0]\n            ==\n            NUM_CLASSES\n        ):\n\n            possible_heads.append(\n                (\n                    key,\n                    tuple(\n                        value.shape\n                    ),\n                )\n            )\n\n    print()\n    print(\n        \"Checkpoint inspection:\"\n    )\n\n    print(\n        f\"  File: {checkpoint_path}\"\n    )\n\n    print(\n        f\"  Candidate 12-class heads:\"\n        f\" {len(possible_heads)}\"\n    )\n\n    for key, shape in (\n        possible_heads[:10]\n    ):\n\n        print(\n            f\"    {key}: {shape}\"\n        )\n\n    return len(\n        possible_heads\n    ) > 0\n\n\n# ============================================================\n# 35. LOAD CHECKPOINT STATE\n# ============================================================\n\ndef load_checkpoint_state(\n    model,\n    checkpoint_path,\n):\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"LOADING LOCAL CHECKPOINT\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        f\"Checkpoint:\"\n    )\n\n    print(\n        f\"  {checkpoint_path}\"\n    )\n\n    checkpoint = torch.load(\n\n        checkpoint_path,\n\n        map_location=\"cpu\",\n\n        weights_only=False,\n    )\n\n    state_dict = None\n\n    if isinstance(\n        checkpoint,\n        dict,\n    ):\n\n        for key in [\n\n            \"state_dict\",\n\n            \"model_state_dict\",\n\n            \"model\",\n\n            \"weights\",\n        ]:\n\n            if key in checkpoint:\n\n                candidate = (\n                    checkpoint[key]\n                )\n\n                if isinstance(\n                    candidate,\n                    dict,\n                ):\n\n                    state_dict = (\n                        candidate\n                    )\n\n                    break\n\n        if state_dict is None:\n\n            if all(\n\n                isinstance(\n                    k,\n                    str,\n                )\n\n                for k in checkpoint.keys()\n\n            ):\n\n                state_dict = checkpoint\n\n    if state_dict is None:\n\n        raise RuntimeError(\n\n            \"Could not locate a state_dict \"\n            \"inside the checkpoint.\"\n        )\n\n    # --------------------------------------------------------\n    # Remove DataParallel prefix.\n    # --------------------------------------------------------\n\n    cleaned = {}\n\n    for key, value in (\n        state_dict.items()\n    ):\n\n        if key.startswith(\n            \"module.\"\n        ):\n\n            key = key[\n                len(\"module.\") :\n            ]\n\n        cleaned[key] = value\n\n    state_dict = cleaned\n\n    # --------------------------------------------------------\n    # Try strict loading first.\n    # --------------------------------------------------------\n\n    try:\n\n        result = model.load_state_dict(\n\n            state_dict,\n\n            strict=True,\n        )\n\n        print(\n            \"Checkpoint loaded with \"\n            \"strict=True.\"\n        )\n\n        return True\n\n    except Exception as exc:\n\n        print()\n        print(\n            \"Strict checkpoint loading failed:\"\n        )\n\n        print(\n            str(exc)[:3000]\n        )\n\n    # --------------------------------------------------------\n    # Try non-strict loading.\n    # --------------------------------------------------------\n\n    try:\n\n        result = model.load_state_dict(\n\n            state_dict,\n\n            strict=False,\n        )\n\n        print()\n        print(\n            \"Checkpoint loaded with \"\n            \"strict=False.\"\n        )\n\n        print(\n            f\"Missing keys: \"\n            f\"{len(result.missing_keys)}\"\n        )\n\n        print(\n            f\"Unexpected keys: \"\n            f\"{len(result.unexpected_keys)}\"\n        )\n\n        return True\n\n    except Exception as exc:\n\n        print()\n        print(\n            \"Non-strict checkpoint loading \"\n            \"also failed:\"\n        )\n\n        print(\n            str(exc)[:3000]\n        )\n\n        return False\n\n\n# ============================================================\n# 36. LOAD MODEL — STRICTLY OFFLINE\n# ============================================================\n\ndef load_model():\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"LOADING OFFLINE MODEL\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    # --------------------------------------------------------\n    # STEP A\n    # Look for an actual trained local checkpoint.\n    # --------------------------------------------------------\n\n    local_files = (\n        find_local_model_files()\n    )\n\n    if local_files:\n\n        print()\n        print(\n            \"Local model files found:\"\n        )\n\n        for path in local_files[:20]:\n\n            print(\n                f\"  {path}\"\n            )\n\n        # ----------------------------------------------------\n        # Prefer explicit MODEL_CHECKPOINT.\n        # ----------------------------------------------------\n\n        ordered = []\n\n        if MODEL_CHECKPOINT is not None:\n\n            explicit = Path(\n                MODEL_CHECKPOINT\n            )\n\n            if explicit in local_files:\n\n                ordered.append(\n                    explicit\n                )\n\n        for path in local_files:\n\n            if path not in ordered:\n\n                ordered.append(\n                    path\n                )\n\n        # ----------------------------------------------------\n        # Try compatible checkpoints.\n        # ----------------------------------------------------\n\n        for checkpoint_path in ordered:\n\n            if not checkpoint_path.is_file():\n\n                continue\n\n            try:\n\n                if not checkpoint_looks_compatible(\n                    checkpoint_path\n                ):\n\n                    continue\n\n                model = (\n                    CheckpointFreeKneeModel(\n\n                        arch=DEFAULT_ARCH,\n\n                        num_classes=NUM_CLASSES,\n\n                        pretrained=False,\n                    )\n                )\n\n                loaded = (\n                    load_checkpoint_state(\n\n                        model,\n\n                        checkpoint_path,\n                    )\n                )\n\n                if loaded:\n\n                    model = model.to(\n                        DEVICE\n                    )\n\n                    model.eval()\n\n                    print()\n                    print(\n                        \"SUCCESS:\"\n                    )\n\n                    print(\n                        \"Using local trained \"\n                        \"RSNA checkpoint.\"\n                    )\n\n                    return model\n\n            except Exception as exc:\n\n                print()\n                print(\n                    \"Checkpoint attempt failed:\"\n                )\n\n                print(\n                    f\"  {checkpoint_path}\"\n                )\n\n                print(\n                    f\"  {exc}\"\n                )\n\n    # --------------------------------------------------------\n    # STEP B\n    #\n    # IMPORTANT:\n    #\n    # We DO NOT call:\n    #\n    # timm.create_model(..., pretrained=True)\n    #\n    # because that can attempt Hugging Face access.\n    #\n    # Instead, we create the architecture locally.\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"No compatible local RSNA checkpoint found.\"\n    )\n\n    print()\n    print(\n        \"Creating CoAtNet with pretrained=False.\"\n    )\n\n    print(\n        \"This guarantees no Hugging Face download.\"\n    )\n\n    model = (\n        CheckpointFreeKneeModel(\n\n            arch=DEFAULT_ARCH,\n\n            num_classes=NUM_CLASSES,\n\n            pretrained=False,\n        )\n    )\n\n    model = model.to(\n        DEVICE\n    )\n\n    model.eval()\n\n    print()\n    print(\n        \"WARNING:\"\n    )\n\n    print(\n        \"The current model has randomly initialized \"\n        \"weights.\"\n    )\n\n    print(\n        \"It is NOT a trained RSNA 12-label model.\"\n    )\n\n    print()\n    print(\n        \"The pipeline is now completely offline.\"\n    )\n\n    return model\n\n\n# ============================================================\n# 37. PREDICT WINDOWS\n# ============================================================\n\n@torch.inference_mode()\ndef predict_windows(\n    model,\n    windows,\n    batch_size=16,\n):\n\n    if len(windows) == 0:\n\n        return np.zeros(\n\n            NUM_CLASSES,\n\n            dtype=np.float32,\n        )\n\n    predictions = []\n\n    for start in range(\n\n        0,\n\n        len(windows),\n\n        batch_size,\n    ):\n\n        batch = torch.from_numpy(\n\n            windows[\n\n                start:\n                start + batch_size\n            ]\n        ).to(\n\n            DEVICE,\n\n            non_blocking=True,\n        )\n\n        try:\n\n            if DEVICE.type == \"cuda\":\n\n                with torch.autocast(\n\n                    device_type=\"cuda\",\n\n                    dtype=torch.float16,\n                ):\n\n                    logits = model(\n                        batch\n                    )\n\n            else:\n\n                logits = model(\n                    batch\n                )\n\n        except Exception:\n\n            logits = model(\n                batch\n            )\n\n        if logits.ndim == 1:\n\n            logits = logits.unsqueeze(\n                0\n            )\n\n        probs = torch.sigmoid(\n\n            logits.float()\n        )\n\n        predictions.append(\n\n            probs.detach()\n            .cpu()\n            .numpy()\n        )\n\n    predictions = np.concatenate(\n\n        predictions,\n\n        axis=0,\n    )\n\n    return predictions.mean(\n\n        axis=0\n\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 38. TTA FUSION\n# ============================================================\n\ndef fuse_tta(\n    predictions,\n    weights,\n):\n\n    predictions = np.asarray(\n\n        predictions,\n\n        dtype=np.float32,\n    )\n\n    weights = np.asarray(\n\n        weights,\n\n        dtype=np.float32,\n    )\n\n    weight_sum = weights.sum()\n\n    if weight_sum <= 0:\n\n        weights = np.ones_like(\n            weights\n        )\n\n        weight_sum = weights.sum()\n\n    weights = (\n\n        weights\n        /\n        weight_sum\n    )\n\n    probabilities = np.clip(\n\n        predictions,\n\n        1e-6,\n\n        1.0 - 1e-6,\n    )\n\n    logits = np.log(\n\n        probabilities\n\n        /\n\n        (\n            1.0\n            -\n            probabilities\n        )\n    )\n\n    fused_logits = np.sum(\n\n        logits\n        *\n        weights[:, None],\n\n        axis=0,\n    )\n\n    fused_probs = (\n\n        1.0\n\n        /\n\n        (\n\n            1.0\n            +\n            np.exp(\n                -fused_logits\n            )\n        )\n    )\n\n    direct_probs = np.sum(\n\n        probabilities\n        *\n        weights[:, None],\n\n        axis=0,\n    )\n\n    fused = (\n\n        0.60\n        *\n        fused_probs\n\n        +\n\n        0.40\n        *\n        direct_probs\n    )\n\n    return np.clip(\n\n        fused,\n\n        0.0,\n\n        1.0,\n\n    ).astype(\n        np.float32\n    )\n\n\n# ============================================================\n# 39. PREDICT ONE STUDY\n# ============================================================\n\ndef predict_study(\n    model,\n    series_lookup,\n):\n\n    volume, valid = (\n        build_volume(\n\n            series_lookup,\n\n            TOTAL_SLICES,\n        )\n    )\n\n    valid_count = int(\n        valid.sum()\n    )\n\n    if valid_count == 0:\n\n        print(\n            \"WARNING: study contains \"\n            \"no valid DICOM slices.\"\n        )\n\n        return np.full(\n\n            NUM_CLASSES,\n\n            0.5,\n\n            dtype=np.float32,\n        )\n\n    tta_predictions = []\n\n    tta_weights = []\n\n    for config in TTA_CONFIGS:\n\n        centers = make_centers(\n\n            valid,\n\n            config[\"k\"],\n\n            config[\"lo\"],\n\n            config[\"hi\"],\n        )\n\n        windows = create_windows(\n\n            volume,\n\n            centers,\n        )\n\n        prediction = (\n            predict_windows(\n\n                model,\n\n                windows,\n            )\n        )\n\n        tta_predictions.append(\n            prediction\n        )\n\n        tta_weights.append(\n            config[\"weight\"]\n        )\n\n        del windows\n\n        if DEVICE.type == \"cuda\":\n\n            torch.cuda.empty_cache()\n\n    result = fuse_tta(\n\n        tta_predictions,\n\n        tta_weights,\n    )\n\n    del volume\n\n    del valid\n\n    return result\n\n\n# ============================================================\n# 40. FINAL RANK CALIBRATION\n# ============================================================\n\ndef final_rank_calibration(\n    predictions,\n):\n\n    predictions = np.asarray(\n\n        predictions,\n\n        dtype=np.float32,\n    ).copy()\n\n    if len(predictions) <= 1:\n\n        return np.clip(\n\n            predictions,\n\n            0.0,\n\n            1.0,\n        )\n\n    for j, label in enumerate(\n        LABELS\n    ):\n\n        values = predictions[\n            :,\n            j,\n        ]\n\n        order = np.argsort(\n\n            np.argsort(\n                values\n            )\n        )\n\n        ranks = (\n\n            order.astype(\n                np.float32\n            )\n\n            /\n\n            max(\n\n                len(values) - 1,\n\n                1,\n            )\n        )\n\n        weight = (\n\n            TARGET_RANK_WEIGHTS.get(\n\n                label,\n\n                0.0,\n            )\n        )\n\n        predictions[\n\n            :,\n\n            j\n\n        ] = (\n\n            (\n\n                1.0\n                -\n                weight\n            )\n\n            *\n\n            values\n\n            +\n\n            weight\n\n            *\n\n            ranks\n        )\n\n    return np.clip(\n\n        predictions,\n\n        0.0,\n\n        1.0,\n    )\n\n\n# ============================================================\n# 41. SAMPLE SUBMISSION\n# ============================================================\n\ndef load_sample_submission(\n    root\n):\n\n    path = (\n\n        root /\n\n        \"sample_submission.csv\"\n    )\n\n    if path.is_file():\n\n        return pd.read_csv(\n            path\n        )\n\n    matches = list(\n\n        root.glob(\n\n            \"**/sample_submission.csv\"\n        )\n    )\n\n    if matches:\n\n        return pd.read_csv(\n            matches[0]\n        )\n\n    return None\n\n\n# ============================================================\n# 42. VALIDATE SUBMISSION\n# ============================================================\n\ndef validate_submission(\n    submission,\n    sample_submission=None,\n):\n\n    if submission.empty:\n\n        raise ValueError(\n            \"Submission is empty.\"\n        )\n\n    id_col = (\n        submission.columns[0]\n    )\n\n    expected_columns = [\n\n        id_col\n\n    ] + LABELS\n\n    missing = [\n\n        col\n\n        for col in expected_columns\n\n        if col not in submission.columns\n    ]\n\n    if missing:\n\n        raise ValueError(\n\n            f\"Missing columns: {missing}\"\n        )\n\n    prediction_values = (\n\n        submission[\n            LABELS\n        ].to_numpy()\n    )\n\n    if not np.isfinite(\n        prediction_values\n    ).all():\n\n        raise ValueError(\n\n            \"Submission contains \"\n            \"NaN or infinite values.\"\n        )\n\n    if (\n\n        prediction_values.min()\n        <\n        0\n\n        or\n\n        prediction_values.max()\n        >\n        1\n\n    ):\n\n        raise ValueError(\n\n            \"Predictions must be \"\n            \"between 0 and 1.\"\n        )\n\n    if sample_submission is not None:\n\n        expected = list(\n\n            sample_submission.columns\n        )\n\n        actual = list(\n\n            submission.columns\n        )\n\n        if expected != actual:\n\n            raise ValueError(\n\n                \"\\nSubmission columns do not \"\n                \"match sample_submission.\\n\"\n\n                f\"Expected: {expected}\\n\"\n\n                f\"Actual:   {actual}\"\n            )\n\n        if (\n\n            len(submission)\n\n            !=\n\n            len(sample_submission)\n\n        ):\n\n            raise ValueError(\n\n                \"Submission row count does \"\n                \"not match sample_submission.\"\n            )\n\n        id_col = expected[0]\n\n        actual_ids = (\n\n            submission[\n                id_col\n            ].astype(str)\n        )\n\n        expected_ids = (\n\n            sample_submission[\n                id_col\n            ].astype(str)\n        )\n\n        if not actual_ids.equals(\n            expected_ids\n        ):\n\n            raise ValueError(\n\n                \"Submission IDs/order do \"\n                \"not match sample_submission.\"\n            )\n\n    print(\n        \"Submission validation: PASSED\"\n    )\n\n\n# ============================================================\n# 43. SANITY CHECK LOOKUP\n# ============================================================\n\ndef sanity_check_lookup(\n    lookup,\n    test_studies,\n):\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"LOOKUP SANITY CHECK\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    total_series = 0\n\n    total_slices = 0\n\n    studies_with_data = 0\n\n    for study_id in test_studies:\n\n        study = lookup.get(\n\n            str(study_id),\n\n            {},\n        )\n\n        series_count = len(\n            study\n        )\n\n        slice_count = sum(\n\n            int(\n\n                s.get(\n\n                    \"image_count\",\n\n                    0,\n                )\n            )\n\n            for s in study.values()\n        )\n\n        total_series += (\n            series_count\n        )\n\n        total_slices += (\n            slice_count\n        )\n\n        if slice_count > 0:\n\n            studies_with_data += 1\n\n        print(\n\n            f\"Study {study_id}: \"\n\n            f\"{series_count} series | \"\n\n            f\"{slice_count} slices\"\n        )\n\n    print()\n\n    print(\n\n        f\"Studies with DICOMs: \"\n\n        f\"{studies_with_data}/\"\n\n        f\"{len(test_studies)}\"\n    )\n\n    print(\n\n        f\"Total series: \"\n\n        f\"{total_series:,}\"\n    )\n\n    print(\n\n        f\"Total DICOM slices: \"\n\n        f\"{total_slices:,}\"\n    )\n\n    if (\n\n        studies_with_data\n\n        ==\n\n        0\n\n    ):\n\n        raise RuntimeError(\n\n            \"\\nNO TEST DICOM DATA FOUND.\\n\"\n\n            \"\\n\"\n\n            \"Check that the competition input \"\n            \"contains test_series/.\\n\"\n        )\n\n    print(\n        \"Lookup sanity check: PASSED\"\n    )\n\n\n# ============================================================\n# 44. BUILD TEST LOOKUP\n# ============================================================\n\ndef build_lookup_for_test(\n    root,\n    test_df,\n    id_col,\n):\n\n    test_series_root = (\n        locate_test_series(\n            root\n        )\n    )\n\n    (\n        test_series_df,\n        test_series_csv,\n    ) = (\n\n        load_test_series_csv(\n            root\n        )\n    )\n\n    test_studies = (\n\n        test_df[\n            id_col\n        ]\n\n        .astype(str)\n\n        .tolist()\n    )\n\n    lookup = (\n\n        create_series_lookup_optimized(\n\n            test_series_root,\n\n            test_studies,\n\n            test_series_df,\n\n            root / \"test.csv\",\n\n            test_series_csv,\n\n            CACHE_FILE,\n        )\n    )\n\n    sanity_check_lookup(\n\n        lookup,\n\n        test_studies,\n    )\n\n    return lookup\n\n\n# ============================================================\n# 45. OFFLINE STATUS CHECK\n# ============================================================\n\ndef print_offline_status():\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"OFFLINE STATUS CHECK\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"Hugging Face offline:\",\n        os.environ.get(\n            \"HF_HUB_OFFLINE\"\n        )\n    )\n\n    print(\n        \"Transformers offline:\",\n        os.environ.get(\n            \"TRANSFORMERS_OFFLINE\"\n        )\n    )\n\n    print(\n        \"Xet disabled:\",\n        os.environ.get(\n            \"HF_HUB_DISABLE_XET\"\n        )\n    )\n\n    print()\n\n    print(\n        \"No internet download will be \"\n        \"attempted by this pipeline.\"\n    )\n\n\n# ============================================================\n# 46. MAIN\n# ============================================================\n\ndef main():\n\n    total_start = time.time()\n\n    # --------------------------------------------------------\n    # OFFLINE STATUS\n    # --------------------------------------------------------\n\n    print_offline_status()\n\n    # --------------------------------------------------------\n    # STEP 1\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 1 — LOCATING DATA\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    root = (\n        locate_competition_root()\n    )\n\n    # --------------------------------------------------------\n    # STEP 2\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 2 — READING TEST CSV\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    (\n        test_df,\n        id_col,\n        test_csv,\n    ) = load_test_csv(\n        root\n    )\n\n    # --------------------------------------------------------\n    # STEP 3\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 3 — BUILDING TEST SERIES LOOKUP\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    lookup = (\n        build_lookup_for_test(\n\n            root,\n\n            test_df,\n\n            id_col,\n        )\n    )\n\n    # --------------------------------------------------------\n    # STEP 4\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 4 — LOADING OFFLINE MODEL\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    model = load_model()\n\n    # --------------------------------------------------------\n    # STEP 5\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 5 — RUNNING OFFLINE INFERENCE\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    predictions = []\n\n    study_ids = (\n\n        test_df[\n            id_col\n        ]\n\n        .astype(str)\n\n        .tolist()\n    )\n\n    inference_start = time.time()\n\n    for (\n\n        index,\n\n        study_id,\n\n    ) in enumerate(\n\n        study_ids,\n\n        start=1,\n    ):\n\n        series_lookup = (\n\n            lookup.get(\n\n                study_id,\n\n                {},\n            )\n        )\n\n        try:\n\n            prediction = (\n\n                predict_study(\n\n                    model,\n\n                    series_lookup,\n                )\n            )\n\n        except Exception as exc:\n\n            print()\n\n            print(\n\n                f\"WARNING: inference failed \"\n                f\"for study {study_id}\"\n            )\n\n            print(\n                f\"Reason: {exc}\"\n            )\n\n            prediction = np.full(\n\n                NUM_CLASSES,\n\n                0.5,\n\n                dtype=np.float32,\n            )\n\n        predictions.append(\n            prediction\n        )\n\n        if (\n\n            index == 1\n\n            or\n\n            index\n            %\n            PROGRESS_EVERY_STUDIES\n            ==\n            0\n\n            or\n\n            index\n            ==\n            len(study_ids)\n\n        ):\n\n            elapsed = (\n\n                time.time()\n                -\n                inference_start\n            )\n\n            rate = (\n\n                index\n                /\n                max(\n                    elapsed,\n                    1e-6,\n                )\n            )\n\n            remaining = (\n\n                len(study_ids)\n                -\n                index\n            )\n\n            eta = (\n\n                remaining\n                /\n                max(\n                    rate,\n                    1e-6,\n                )\n            )\n\n            print(\n\n                f\"[{index:,}/\"\n                f\"{len(study_ids):,}] \"\n\n                f\"studies | \"\n\n                f\"{rate:.3f} studies/s | \"\n\n                f\"ETA: \"\n                f\"{eta / 60:.2f} min\"\n            )\n\n        gc.collect()\n\n    predictions = np.vstack(\n        predictions\n    )\n\n    # --------------------------------------------------------\n    # STEP 6\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 6 — FINAL RANK CALIBRATION\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    calibrated = (\n\n        final_rank_calibration(\n\n            predictions\n        )\n    )\n\n    # --------------------------------------------------------\n    # STEP 7\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"STEP 7 — CREATING OFFLINE SUBMISSION\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    submission = pd.DataFrame(\n\n        calibrated,\n\n        columns=LABELS,\n    )\n\n    submission.insert(\n\n        0,\n\n        id_col,\n\n        test_df[\n            id_col\n        ].values,\n    )\n\n    sample_submission = (\n\n        load_sample_submission(\n            root\n        )\n    )\n\n    validate_submission(\n\n        submission,\n\n        sample_submission,\n    )\n\n    submission.to_csv(\n\n        OUTPUT_FILE,\n\n        index=False,\n    )\n\n    print()\n    print(\n        \"Submission saved:\"\n    )\n\n    print(\n        f\"  {OUTPUT_FILE}\"\n    )\n\n    print()\n\n    print(\n        submission.head()\n    )\n\n    # --------------------------------------------------------\n    # Prediction statistics\n    # --------------------------------------------------------\n\n    print()\n    print(\n        \"Prediction statistics:\"\n    )\n\n    print()\n\n    stats = pd.DataFrame({\n\n        \"label\":\n            LABELS,\n\n        \"min\":\n            calibrated.min(\n                axis=0\n            ),\n\n        \"mean\":\n            calibrated.mean(\n                axis=0\n            ),\n\n        \"max\":\n            calibrated.max(\n                axis=0\n            ),\n    })\n\n    print(\n\n        stats.to_string(\n            index=False\n        )\n    )\n\n    # --------------------------------------------------------\n    # STEP 8\n    # --------------------------------------------------------\n\n    total_elapsed = (\n\n        time.time()\n        -\n        total_start\n    )\n\n    print()\n    print(\n        \"=\" * 70\n    )\n\n    print(\n        \"PIPELINE COMPLETE\"\n    )\n\n    print(\n        \"=\" * 70\n    )\n\n    print(\n\n        f\"Studies: \"\n        f\"{len(submission):,}\"\n    )\n\n    print(\n\n        f\"Labels: \"\n        f\"{NUM_CLASSES}\"\n    )\n\n    print()\n\n    print(\n        \"Output:\"\n    )\n\n    print(\n        f\"  {OUTPUT_FILE}\"\n    )\n\n    print()\n\n    print(\n        \"Total runtime:\"\n    )\n\n    print(\n        f\"  {total_elapsed / 60:.2f} minutes\"\n    )\n\n    print()\n\n    print(\n        \"OFFLINE SUBMISSION READY.\"\n    )\n\n    return submission\n\n\n# ============================================================\n# 47. RUN\n# ============================================================\n\nif __name__ == \"__main__\":\n\n    submission = main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T08:01:49.952926Z","iopub.execute_input":"2026-09-14T08:01:49.953244Z","iopub.status.idle":"2026-09-14T08:08:08.035595Z","shell.execute_reply.started":"2026-09-14T08:01:49.953218Z","shell.execute_reply":"2026-09-14T08:08:08.0347Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <div style=\"color:white;display:inline-block;border-radius:4px;background-color:#065535 ;font-family:Nexa;overflow:hidden\"><p style=\"padding:8px;color:white;overflow:hidden;font-size:85%;letter-spacing:0.5px;margin:0;border: 6px groove #e4c155;\"><b> </b>Final Image Visualization</p></div>","metadata":{}},{"cell_type":"code","source":"#  FINAL IMAGE VISUALIZATIONS\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\n\n# ============================================================\n# BASIC VALIDATION\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL IMAGE VISUALIZATIONS\")\nprint(\"=\" * 70)\n\n# ------------------------------------------------------------\n# Validate submission\n# ------------------------------------------------------------\n\nif \"submission\" not in globals():\n    raise NameError(\n        \"The `submission` DataFrame does not exist. \"\n        \"Please run the prediction/submission cells first.\"\n    )\n\nif \"LABELS\" not in globals():\n    raise NameError(\n        \"The `LABELS` variable is not defined. \"\n        \"Please run the label-definition cell first.\"\n    )\n\nif len(submission) == 0:\n    raise ValueError(\n        \"The submission DataFrame is empty.\"\n    )\n\n\n# ============================================================\n# AUTOMATICALLY FIND ID COLUMN\n# ============================================================\n\n# Common ID column names used in RSNA-style datasets.\npossible_id_columns = [\n    \"StudyInstanceUID\",\n    \"study_id\",\n    \"studyId\",\n    \"StudyID\",\n    \"id\",\n    \"ID\",\n    \"patient_id\",\n    \"PatientID\",\n]\n\nid_col = None\n\nfor candidate in possible_id_columns:\n\n    if candidate in submission.columns:\n\n        id_col = candidate\n        break\n\n\n# ------------------------------------------------------------\n# Fallback:\n# use the first column that is NOT a prediction label\n# ------------------------------------------------------------\n\nif id_col is None:\n\n    non_label_columns = [\n        col\n        for col in submission.columns\n        if col not in LABELS\n    ]\n\n    if len(non_label_columns) > 0:\n\n        id_col = non_label_columns[0]\n\n\n# ------------------------------------------------------------\n# Stop with a clear message if no ID column exists\n# ------------------------------------------------------------\n\nif id_col is None:\n\n    raise ValueError(\n        \"Could not identify the submission ID column.\\n\"\n        f\"Available columns: {list(submission.columns)}\"\n    )\n\n\nprint()\nprint(f\"Submission rows : {len(submission):,}\")\nprint(f\"ID column       : {id_col}\")\nprint(f\"Prediction cols : {len(LABELS)}\")\n\n\n# ============================================================\n# VALIDATE LABEL COLUMNS\n# ============================================================\n\nmissing_labels = [\n    label\n    for label in LABELS\n    if label not in submission.columns\n]\n\nif len(missing_labels) > 0:\n\n    raise ValueError(\n        \"The following LABELS are missing from submission:\\n\"\n        f\"{missing_labels}\\n\\n\"\n        f\"Available columns:\\n\"\n        f\"{list(submission.columns)}\"\n    )\n\n\n# ============================================================\n# OUTPUT DIRECTORY\n# ============================================================\n\nOUTPUT_DIR = \"/kaggle/working\"\n\nos.makedirs(\n    OUTPUT_DIR,\n    exist_ok=True,\n)\n\n\n# ============================================================\n# VISUALIZATION 1\n# Prediction Distribution\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"VISUALIZATION 1 — PREDICTION DISTRIBUTION\")\nprint(\"=\" * 70)\n\n\nplt.figure(\n    figsize=(12, 6)\n)\n\n\nfor label in LABELS:\n\n    values = pd.to_numeric(\n        submission[label],\n        errors=\"coerce\",\n    ).dropna()\n\n    plt.hist(\n        values,\n        bins=30,\n        alpha=0.35,\n        label=label,\n    )\n\n\nplt.xlabel(\n    \"Predicted probability\"\n)\n\nplt.ylabel(\n    \"Number of studies\"\n)\n\nplt.title(\n    \"Distribution of Predicted Knee Abnormality Probabilities\"\n)\n\nplt.legend(\n    bbox_to_anchor=(1.02, 1),\n    loc=\"upper left\",\n    fontsize=9,\n)\n\nplt.tight_layout()\n\n\nviz1_path = os.path.join(\n    OUTPUT_DIR,\n    \"prediction_distribution.png\",\n)\n\n\nplt.savefig(\n    viz1_path,\n    dpi=180,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\nplt.close()\n\n\nprint(\n    f\"Saved: {viz1_path}\"\n)\n\n\n# ============================================================\n# VISUALIZATION 2\n# Prediction Heatmap\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"VISUALIZATION 2 — PREDICTION HEATMAP\")\nprint(\"=\" * 70)\n\n\n# ------------------------------------------------------------\n# Number of studies to display\n# ------------------------------------------------------------\n\nn_show = min(\n    30,\n    len(submission),\n)\n\n\n# ------------------------------------------------------------\n# Extract prediction values\n# ------------------------------------------------------------\n\nheatmap_data = (\n    submission[LABELS]\n    .iloc[:n_show]\n    .apply(\n        pd.to_numeric,\n        errors=\"coerce\",\n    )\n    .fillna(0.0)\n    .values\n)\n\n\n# ------------------------------------------------------------\n# Extract study IDs\n# ------------------------------------------------------------\n\nstudy_ids = (\n    submission[id_col]\n    .iloc[:n_show]\n    .astype(str)\n    .tolist()\n)\n\n\n# ============================================================\n# CREATE HEATMAP\n# ============================================================\n\nplt.figure(\n    figsize=(14, 8)\n)\n\n\nimage = plt.imshow(\n    heatmap_data,\n    aspect=\"auto\",\n    interpolation=\"nearest\",\n)\n\n\nplt.colorbar(\n    image,\n    label=\"Predicted probability\",\n)\n\n\nplt.xticks(\n    range(len(LABELS)),\n    LABELS,\n    rotation=45,\n    ha=\"right\",\n)\n\n\nplt.yticks(\n    range(n_show),\n    study_ids,\n)\n\n\nplt.xlabel(\n    \"Knee abnormality\"\n)\n\nplt.ylabel(\n    \"Study ID\"\n)\n\nplt.title(\n    \"Knee Abnormality Prediction Heatmap\"\n)\n\nplt.tight_layout()\n\n\nviz2_path = os.path.join(\n    OUTPUT_DIR,\n    \"prediction_heatmap.png\",\n)\n\n\nplt.savefig(\n    viz2_path,\n    dpi=180,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\nplt.close()\n\n\nprint(\n    f\"Saved: {viz2_path}\"\n)\n\n\n# ============================================================\n# VISUALIZATION 3\n# Mean Prediction by Abnormality\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"VISUALIZATION 3 — MEAN PREDICTION BY LABEL\")\nprint(\"=\" * 70)\n\n\nmean_predictions = (\n    submission[LABELS]\n    .apply(\n        pd.to_numeric,\n        errors=\"coerce\",\n    )\n    .mean()\n    .sort_values(\n        ascending=False\n    )\n)\n\n\nplt.figure(\n    figsize=(12, 6)\n)\n\n\nbars = plt.bar(\n    mean_predictions.index,\n    mean_predictions.values,\n)\n\n\nplt.xlabel(\n    \"Knee abnormality\"\n)\n\nplt.ylabel(\n    \"Mean predicted probability\"\n)\n\nplt.title(\n    \"Mean Predicted Probability by Knee Abnormality\"\n)\n\nplt.xticks(\n    rotation=45,\n    ha=\"right\",\n)\n\n\n# ------------------------------------------------------------\n# Print values above bars\n# ------------------------------------------------------------\n\nfor bar, value in zip(\n    bars,\n    mean_predictions.values,\n):\n\n    plt.text(\n        bar.get_x()\n        + bar.get_width() / 2,\n        bar.get_height(),\n        f\"{value:.3f}\",\n        ha=\"center\",\n        va=\"bottom\",\n        fontsize=9,\n    )\n\n\nplt.tight_layout()\n\n\nviz3_path = os.path.join(\n    OUTPUT_DIR,\n    \"mean_prediction_by_label.png\",\n)\n\n\nplt.savefig(\n    viz3_path,\n    dpi=180,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\nplt.close()\n\n\nprint(\n    f\"Saved: {viz3_path}\"\n)\n\n\n# ============================================================\n# FINAL OUTPUT SUMMARY\n# ============================================================\n\nprint()\nprint(\"=\" * 70)\nprint(\"FINAL OUTPUTS\")\nprint(\"=\" * 70)\n\n\nprint()\nprint(\"1. Submission:\")\nprint(f\"   {OUTPUT_FILE}\")\n\n\nprint()\nprint(\"2. Lookup cache:\")\nprint(f\"   {CACHE_FILE}\")\n\n\nprint()\nprint(\"3. Prediction distribution:\")\nprint(f\"   {viz1_path}\")\n\n\nprint()\nprint(\"4. Prediction heatmap:\")\nprint(f\"   {viz2_path}\")\n\n\nprint()\nprint(\"5. Mean prediction by label:\")\nprint(f\"   {viz3_path}\")\n\n\nprint()\nprint(\"=\" * 70)\nprint(\"VISUALIZATION COMPLETE\")\nprint(\"=\" * 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-14T08:28:25.569786Z","iopub.execute_input":"2026-09-14T08:28:25.570645Z","iopub.status.idle":"2026-09-14T08:28:28.070967Z","shell.execute_reply.started":"2026-09-14T08:28:25.570568Z","shell.execute_reply":"2026-09-14T08:28:28.070053Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}