{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\n\nos.environ[\"HF_HUB_OFFLINE\"] = \"1\"\nos.environ[\"TRANSFORMERS_OFFLINE\"] = \"1\"\nos.environ[\"TIMM_HUB_OFFLINE\"] = \"1\"\nos.environ[\"HF_DATASETS_OFFLINE\"] = \"1\"\nos.environ[\"TORCH_HOME\"] = os.environ.get(\"TORCH_HOME\", \"./torch_cache\")\n\n_EXPECTED_COLS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\",\n    \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\n_REPORT_COL_HINTS = (\"report\", \"reports\", \"findings\", \"impression\", \"report_text\")\n\n_LISTDIR_BUDGET = [400]\n\n\ndef _walk_dirs(root, max_depth=5):\n    root = os.path.abspath(root)\n    base_depth = root.rstrip(\"/\\\\\").count(os.sep)\n    for dirpath, dirnames, _files in os.walk(root):\n        depth = dirpath.rstrip(\"/\\\\\").count(os.sep) - base_depth\n        if depth >= max_depth:\n            dirnames[:] = []\n        yield dirpath, dirnames\n\n\ndef _list_children(path):\n    if _LISTDIR_BUDGET[0] <= 0:\n        return []\n    _LISTDIR_BUDGET[0] -= 1\n    try:\n        return sorted(os.listdir(path))\n    except OSError:\n        return []\n\n\ndef _classify_dirname(name):\n    nl = name.strip().lower()\n    if \"test_report\" in nl or nl in (\"reports_test\", \"test-reports\"):\n        return \"test_reports\"\n    if \"train_report\" in nl or nl in (\"reports_train\", \"train-reports\") or \"report\" in nl:\n        return \"train_reports\"\n    vol_like = (\"series\" in nl or \"volume\" in nl or \"image\" in nl\n                or nl in (\"dicom\", \"studies\"))\n    if not vol_like:\n        return None\n    if \"test\" in nl:\n        return \"test_volumes\"\n    return \"train_volumes\"\n\n\ndef _csv_header(path):\n    try:\n        import csv\n        with open(path, newline=\"\", encoding=\"utf-8\", errors=\"ignore\") as f:\n            return [h.strip() for h in next(csv.reader(f), [])]\n    except Exception:\n        return []\n\n\ndef _score_train_csv(path):\n    header = _csv_header(path)\n    return sum(1 for c in _EXPECTED_COLS if c in header)\n\n\ndef _report_column_from_header(header):\n    low = [h.lower() for h in header]\n    for hint in _REPORT_COL_HINTS:\n        if hint in low:\n            return header[low.index(hint)]\n    for h in header:\n        if \"report\" in h.lower():\n            return h\n    return None\n\n\ndef _quick_scan(root):\n    \"\"\"Cheap name-based scan (depth <= 3) that respects a listdir budget, since\n    Kaggle mounts /kaggle/input on a slow FUSE filesystem.\"\"\"\n    train_hits = []\n    named = {}\n    for a in _list_children(root):\n        p1 = os.path.join(root, a)\n        p_csv = os.path.join(p1, \"train.csv\")\n        if os.path.isfile(p_csv):\n            train_hits.append((_score_train_csv(p_csv), p_csv))\n        for b in _list_children(p1):\n            p2 = os.path.join(p1, b)\n            tag = _classify_dirname(b)\n            if tag and tag not in named:\n                named[tag] = p2\n            p_csv = os.path.join(p2, \"train.csv\")\n            if os.path.isfile(p_csv):\n                train_hits.append((_score_train_csv(p_csv), p_csv))\n            for c in _list_children(p2):\n                p3 = os.path.join(p2, c)\n                tag = _classify_dirname(c)\n                if tag and tag not in named:\n                    named[tag] = p3\n                p_csv = os.path.join(p3, \"train.csv\")\n                if os.path.isfile(p_csv):\n                    train_hits.append((_score_train_csv(p_csv), p_csv))\n            if _LISTDIR_BUDGET[0] <= 0:\n                break\n        if _LISTDIR_BUDGET[0] <= 0:\n            break\n    return train_hits, named\n\n\ndef _deep_scan(root, max_dirs=5000, max_depth=5):\n    \"\"\"Bounded fallback used only when the quick scan finds nothing.\"\"\"\n    train_hits, named = [], {}\n    seen = 0\n    for dirpath, dirnames in _walk_dirs(root, max_depth):\n        seen += 1\n        if seen > max_dirs:\n            break\n        p_csv = os.path.join(dirpath, \"train.csv\")\n        if os.path.isfile(p_csv):\n            train_hits.append((_score_train_csv(p_csv), p_csv))\n        for d in dirnames:\n            tag = _classify_dirname(d)\n            if tag and tag not in named:\n                named[tag] = os.path.join(dirpath, d)\n    return train_hits, named\n\n\ndef _resolve_paths():\n    \"\"\"Locate the competition data anywhere under /kaggle/input. Kaggle may nest\n    inputs as /kaggle/input/competitions/<name>/ or /kaggle/input/datasets/<slug>/,\n    and reports may also live as a text column inside train.csv.\"\"\"\n    explicit = \"/kaggle/input/rsna-knee-abnormality-detection\"\n    root = \"/kaggle/input\"\n    _LISTDIR_BUDGET[0] = 400\n\n    train_csv, report_col = None, None\n    named = {}\n\n    if os.path.isfile(os.path.join(explicit, \"train.csv\")):\n        train_csv = os.path.join(explicit, \"train.csv\")\n    elif os.path.isdir(root):\n        hits, named = _quick_scan(root)\n        if not hits:\n            hits, deep_named = _deep_scan(root)\n            for k, v in deep_named.items():\n                named.setdefault(k, v)\n        if hits:\n            hits.sort(reverse=True)\n            train_csv = hits[0][1]\n\n    if train_csv is None:\n        base = explicit\n    else:\n        base = os.path.dirname(train_csv)\n        header = _csv_header(train_csv)\n        report_col = _report_column_from_header(header)\n\n    test_csv = os.path.join(base, \"test.csv\")\n    if not os.path.isfile(test_csv):\n        found_test = None\n        if os.path.isdir(root):\n            for lvl1 in _list_children(root):\n                d1 = os.path.join(root, lvl1)\n                for lvl2 in _list_children(d1):\n                    p = os.path.join(d1, lvl2, \"test.csv\")\n                    if os.path.isfile(p):\n                        found_test = p\n                        break\n                if found_test:\n                    break\n        test_csv = found_test or test_csv\n\n    reports = os.path.join(base, \"train_reports\")\n    if not os.path.isdir(reports):\n        reports = named.get(\"train_reports\") or reports\n    test_reports = os.path.join(base, \"test_reports\")\n    if not os.path.isdir(test_reports):\n        test_reports = named.get(\"test_reports\") or test_reports\n\n    volumes = None\n    for cand in (\"train_series\", \"knee_volumes\", \"train_volumes\", \"train_images\",\n                 \"volumes\", \"series\", \"dicom\"):\n        p = os.path.join(base, cand)\n        if os.path.isdir(p):\n            volumes = p\n            break\n    if volumes is None:\n        volumes = named.get(\"train_volumes\") or os.path.join(base, \"knee_volumes\")\n\n    test_volumes = None\n    for cand in (\"test_series\", \"test_volumes\", \"test_images\"):\n        p = os.path.join(base, cand)\n        if os.path.isdir(p):\n            test_volumes = p\n            break\n    if test_volumes is None:\n        test_volumes = named.get(\"test_volumes\")\n    if test_volumes is None:\n        test_volumes = volumes\n\n    print(f\"[config] BASE_PATH : {base}\")\n    print(f\"[config] train.csv : {train_csv if train_csv else 'NOT FOUND'}\"\n          + (f\" ({_score_train_csv(train_csv)}/{len(_EXPECTED_COLS)} cols)\"\n             if train_csv else \"\"))\n    print(f\"[config] test.csv  : {test_csv if os.path.isfile(test_csv) else 'NOT FOUND'}\")\n    if os.path.isdir(reports):\n        print(f\"[config] reports   : {reports}\")\n    elif report_col:\n        print(f\"[config] reports   : will read column '{report_col}' from train.csv\")\n    else:\n        print(\"[config] reports   : NOT FOUND (text modality disabled)\")\n    print(f\"[config] volumes   : {volumes if os.path.isdir(volumes) else 'NOT FOUND (vision disabled)'}\")\n    print(f\"[config] test vols : {test_volumes if os.path.isdir(test_volumes) else 'NOT FOUND (test uses train root / noise)'}\")\n    if not os.path.isdir(volumes) and not os.path.isdir(reports) and not report_col:\n        print(\"[config] ============================================================\")\n        print(\"[config] WARNING: neither images nor report text were found.\")\n        print(\"[config] The model can then only learn from label statistics\")\n        print(\"[config] and the score will stay near 0.5 (random).\")\n        print(\"[config] Attach the image/report datasets via '+ Add Input', or\")\n        print(\"[config] set REPORTS_DIR / VOLUME_ROOT manually in this file.\")\n        print(\"[config] ============================================================\")\n    return base, train_csv, test_csv, reports, test_reports, volumes, test_volumes\n\n\nBASE_PATH, TRAIN_CSV, TEST_CSV, REPORTS_DIR, TEST_REPORTS_DIR, VOLUME_ROOT, VOLUME_TEST_ROOT = _resolve_paths()\n\nSAMPLE_SUB = os.path.join(BASE_PATH, \"sample_submission.csv\")\n\nOUTPUT_DIR = \"./output\"\nMODEL_DIR = os.path.join(OUTPUT_DIR, \"weights\")\nos.makedirs(MODEL_DIR, exist_ok=True)\n\nWORKING_DIR = \"/kaggle/working\" if os.path.isdir(\"/kaggle/working\") \\\n    else os.path.abspath(\".\")\nSUBMISSION_PATH = os.path.join(WORKING_DIR, \"submission.csv\")\nBEST_MODEL_PATH = os.path.join(WORKING_DIR, \"best_auc_model.pth\")\nCLEAN_WORKING_ON_START = True\n\nTARGET_COLS = [\n    \"ACL\", \"MCL\", \"Medial Meniscus\", \"Lateral Meniscus\",\n    \"Medial OA\", \"Lateral OA\", \"PF OA\",\n    \"Effusion\", \"Synovitis\", \"Baker's\", \"Contusion\", \"Fracture\",\n]\n\nNUM_TARGETS = len(TARGET_COLS)\n\n\nSEQ_COLS = [\"Ax T2\", \"Sag T1\", \"Sag T2 STIR\", \"Sag T2 FSE\", \"Cor T1\", \"Cor T2 STIR\"]\n\n\n\ndef _pick_cache_root():\n    for cand in (\"/tmp\", \"/kaggle/temp\"):\n        if os.path.isdir(cand):\n            return os.path.join(cand, \"vol_cache\")\n    return os.path.join(os.path.abspath(\".\"), \"vol_cache\")\n\n\nVOLUME_CACHE_DIR = _pick_cache_root()\nUSE_VOLUME_CACHE = True\n\n\nCLEAN_CACHE_ON_START = False\n\n\nVOLUME_CACHE_HW = 224\n\nMAX_SLICES = 16\nVOLUME_CACHE_MAX_SLICES = min(16, MAX_SLICES)\n\n\nVOLUME_CACHE_DTYPE = \"uint8\"\n\n\nVOLUME_CACHE_BUDGET_GB = 8.0\nVOLUME_CACHE_WORKERS = 4\n\nVOLUME_CACHE_MIN_PER_SERIES = 4\n\n\n\n\n\nSEED = 42\nFOLDS = 5\n\nEPOCHS = 6\nLR = 2e-4\nLR_BACKBONE = 5e-5\nWD = 1e-4\nWARMUP_RATIO = 0.05\nGRAD_CLIP = 5.0\nLABEL_SMOOTH = 0.01\n\n\nMIXUP_ALPHA = 0.3\nMIXUP_PROB = 0.35\nCUTMIX_ALPHA = 1.0\nCUTMIX_PROB = 0.25\n\n\n\nLOSS = \"asl\"\nASL_GAMMA_NEG = 4.0\nASL_GAMMA_POS = 0.0\nASL_CLIP = 0.05\nFOCAL_GAMMA = 2.0\nPOS_WEIGHT_MAX = 10.0\n\nUSE_AMP = True\n\nAMP_DTYPE = \"float16\"\nEMA_DECAY = 0.995\n\n\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nNUM_GPUS = torch.cuda.device_count()\n\n\nDATA_PARALLEL = False\nCHANNELS_LAST = True\nGRAD_CHECKPOINT = False\n\n\ndef _auto_workers():\n    cpu = os.cpu_count() or 2\n    return max(1, min(2, cpu - 1))\n\n\nNUM_WORKERS = _auto_workers()\nPIN_MEMORY = True\nPREFETCH_FACTOR = 2\nPERSISTENT_WORKERS = False\nGC_EVERY_BATCHES = 40\nHANG_TIMEOUT_S = 600\n\nSHARING_STRATEGY = \"file_system\"\n\n\n\nMODEL_NAME = \"efficientnet_b0\"\n\nBATCH_SIZE = 4\nACCUM_STEPS = 4\nIMG_SIZE = 224\nN_SEQ_IMAGES = 16\nLOG_EVERY_STEPS = 10\n\nIN_CHANS = 1\n\n\nBACKBONE_PRESETS = {\n    \"efficientnet_b0\": dict(IMG_SIZE=224, BATCH_SIZE=4, N_SEQ_IMAGES=16,\n                            GRAD_CHECKPOINT=False),\n    \"efficientnet_b2\": dict(IMG_SIZE=288, BATCH_SIZE=4, N_SEQ_IMAGES=24,\n                            GRAD_CHECKPOINT=False),\n    \"convnext_small\":  dict(IMG_SIZE=224, BATCH_SIZE=4, N_SEQ_IMAGES=20,\n                            GRAD_CHECKPOINT=True),\n    \"swin_tiny_patch4_window7_224\": dict(IMG_SIZE=224, BATCH_SIZE=4,\n                                         N_SEQ_IMAGES=20, GRAD_CHECKPOINT=True),\n}\n\nAPPLY_BACKBONE_PRESET = True\nif APPLY_BACKBONE_PRESET and MODEL_NAME in BACKBONE_PRESETS:\n    for _k, _v in BACKBONE_PRESETS[MODEL_NAME].items():\n        globals()[_k] = _v\n    del _k, _v\n\nN_SEQ_IMAGES = min(N_SEQ_IMAGES, MAX_SLICES)\nVOLUME_CACHE_MAX_SLICES = min(VOLUME_CACHE_MAX_SLICES, MAX_SLICES)\n\n\nVOLUME_CACHE_HW = max(VOLUME_CACHE_HW, IMG_SIZE)\n\nPRETRAINED_WEIGHTS = \"\"\nENCODER_FEAT = 1280\n\n\n\nSLICE_STRATEGY = \"saliency\"\nSLICE_JITTER = 0.25\nSALIENCY_TEMP = 0.5\n\n\n\nAUG_HFLIP_P = 0.5\nAUG_SSR_P = 0.7\nAUG_SHIFT = 0.06\nAUG_SCALE = 0.10\nAUG_ROTATE = 12.0\nAUG_BRIGHTNESS_P = 0.5\nAUG_GAMMA_P = 0.3\nAUG_NOISE_P = 0.2\nAUG_SLICE_DROPOUT_P = 0.15\nTEXT_WORD_DROPOUT = 0.10\n\n\nTEXT_VOCAB = 30522\nTEXT_MAX_LEN = 256\nTEXT_HIDDEN = 384\nTEXT_MIN_DF = 2\n\nFUSION_HIDDEN = 768\nFUSION_DIM = 512\nFUSION_HEADS = 8\nDROPOUT = 0.3\n\n\nSLICE_TRANSFORMER_LAYERS = 2\nSLICE_TRANSFORMER_HEADS = 8\n\n\nMODALITY_DROPOUT_TEXT = 0.25\nMODALITY_DROPOUT_IMAGE = 0.05\n\nFOLD_VOCAB = True\n\n\n\nTTA = True\n\nTTA_OPS = [\"none\", \"hflip\", \"vflip\", \"gamma\"]\n\n\nENSEMBLE = \"rank\"\n\n\nVOLUME_CACHE = 3\nVOLUME_MAX_HW = VOLUME_CACHE_HW\nVOLUME_MAX_HW_EVEN = VOLUME_CACHE_HW\n\n\nif __name__ == \"__main__\":\n    import sys as _sys\n    import types as _types\n    _cell_module = _types.ModuleType(\"config\")\n    _cell_module.__dict__.update(globals())\n    _sys.modules[\"config\"] = _cell_module\n\nif \"VOLUME_TEST_ROOT\" not in dir():\n    VOLUME_TEST_ROOT = VOLUME_ROOT\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport glob\nimport json\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\ntry:\n    import config\nexcept ModuleNotFoundError:\n    import sys as _sys\n    import types as _types\n    config = _types.ModuleType(\"config\")\n    config.__dict__.update(globals())\n    _sys.modules[\"config\"] = config\n\n\ntry:\n    import cv2\n\n    cv2.setNumThreads(0)\n    _HAS_CV2 = True\nexcept Exception:\n    _HAS_CV2 = False\n\n\n_CV2_RESIZE_MAX_CHANNELS = 4\n\n\ndef _as_cv_resizable(arr):\n    \"\"\"Coerce a single plane into a layout cv2.resize accepts: 2D (H, W) or up to\n    4-channel (H, W, C), C-contiguous with a valid depth.\"\"\"\n    arr = np.ascontiguousarray(np.asarray(arr, dtype=np.float32))\n    if arr.ndim == 3:\n        if arr.shape[-1] > 4:\n            arr = np.ascontiguousarray(arr.mean(axis=-1))\n        elif arr.shape[-1] == 1:\n            arr = arr[..., 0]\n    if arr.ndim != 2 and arr.ndim != 3:\n        arr = np.ascontiguousarray(arr.reshape(arr.shape[0], -1))\n    return arr\n\n\ndef _resize2d(arr, h, w):\n    \"\"\"Resize a single 2D float32 plane.\"\"\"\n    if arr.shape[0] == h and arr.shape[1] == w:\n        return arr\n    if _HAS_CV2:\n        plane = _as_cv_resizable(arr)\n        interp = cv2.INTER_AREA if (plane.shape[0] > h or plane.shape[1] > w) \\\n            else cv2.INTER_LINEAR\n        out = cv2.resize(plane, (w, h), interpolation=interp)\n        if out.ndim == 2:\n            out = out[:, :, None]\n        if plane.ndim == 3:\n            return out[:, :, :plane.shape[-1]]\n        return np.ascontiguousarray(out[:, :, 0])\n    from skimage.transform import resize as _sk_resize\n    return _sk_resize(arr, (h, w), order=1, anti_aliasing=True,\n                      preserve_range=True).astype(np.float32)\n\n\ndef resize_volume(vol, h, w):\n    \"\"\"Resize every slice of a (D, H, W) volume, batching channels through cv2.\n\n    cv2.resize only accepts <= 4 channels, so the batched (H, W, D) warp is used\n    only for small stacks; larger/multi-slice volumes fall back to resampling\n    each plane individually to avoid the \"cn <= 4\" assertion failure.\"\"\"\n    vol = np.asarray(vol, dtype=np.float32)\n    if vol.ndim != 3:\n        raise ValueError(f\"resize_volume expects (D, H, W), got shape {vol.shape}\")\n    if vol.shape[1] == h and vol.shape[2] == w:\n        return vol\n    d = vol.shape[0]\n    if _HAS_CV2 and 1 <= d <= _CV2_RESIZE_MAX_CHANNELS:\n\n        stacked = np.ascontiguousarray(vol.transpose(1, 2, 0))\n        interp = cv2.INTER_AREA if (vol.shape[1] > h or vol.shape[2] > w) \\\n            else cv2.INTER_LINEAR\n        out = cv2.resize(stacked, (w, h), interpolation=interp)\n        if out.ndim == 2:\n            out = out[:, :, None]\n        return np.ascontiguousarray(out.transpose(2, 0, 1), dtype=np.float32)\n    out = np.empty((d, h, w), dtype=np.float32)\n    for i in range(d):\n        out[i] = _resize2d(vol[i], h, w)\n    return out\n\n\ndef _input_tree(limit=60, max_depth=3):\n    root = \"/kaggle/input\"\n    if not os.path.isdir(root):\n        return \"  (not on Kaggle; /kaggle/input does not exist)\"\n    lines = []\n    base_depth = os.path.abspath(root).rstrip(\"/\\\\\").count(os.sep)\n    for dirpath, dirnames, files in os.walk(root):\n        depth = dirpath.rstrip(\"/\\\\\").count(os.sep) - base_depth\n        if depth >= max_depth:\n            dirnames[:] = []\n        if len(lines) >= limit:\n            lines.append(\"  ...\")\n            break\n        lines.append(f\"  {'  ' * depth}{os.path.basename(dirpath) or dirpath}/ \"\n                     f\"({len(files)} files)\")\n    return \"\\n\".join(lines)\n\n\ndef _require_file(path, hint):\n    if not os.path.isfile(path):\n        raise FileNotFoundError(\n            f\"No such file: {path}\\n{hint}\\n\"\n            f\"BASE_PATH is set to: {config.BASE_PATH}\\n\"\n            f\"Attached input tree:\\n{_input_tree()}\"\n        )\n    return path\n\n\ndef load_train_df():\n    _require_file(config.TRAIN_CSV,\n                  \"Train data was not found. Did you attach the competition dataset?\\n\"\n                  \"  In the notebook: right sidebar -> '+ Add Input' ->\\n\"\n                  \"  search 'rsna-knee-abnormality-detection' -> click '+'.\")\n    df = pd.read_csv(config.TRAIN_CSV)\n    df = df.fillna(0)\n    return df\n\n\ndef load_test_df():\n    _require_file(config.TEST_CSV,\n                  \"Test data was not found. Verify the competition dataset is attached \"\n                  \"and contains test.csv.\")\n    return pd.read_csv(config.TEST_CSV)\n\n\ndef get_target_columns(df):\n    return [c for c in config.TARGET_COLS if c in df.columns]\n\n\ndef uid_column(df):\n    \"\"\"Study-id column name, tolerating datasets that do not use the canonical name.\"\"\"\n    if \"StudyInstanceUID\" in df.columns:\n        return \"StudyInstanceUID\"\n    for c in df.columns:\n        cl = c.lower()\n        if \"studyinstance\" in cl or cl in (\"study_id\", \"studyid\", \"study\", \"uid\", \"id\"):\n            return c\n    return df.columns[0]\n\n\ndef clean_text(raw, max_len=4000):\n    if not isinstance(raw, str):\n        return \"\"\n    t = re.sub(r\"\\r?\\n\", \" \", raw)\n    t = re.sub(r\"\\[(?:No|an?) impossible to assess[^\\]]*\\]\", \"<UNK>\", t, flags=re.I)\n    t = re.sub(r\"[^0-9a-zA-Z',.:()\\- ]+\", \" \", t)\n    t = re.sub(r\"\\s+\", \" \", t).strip().lower()\n    return t[:max_len]\n\n\ndef load_report_text(report_path):\n    for ext in (\".txt\", \".md\", \".csv\"):\n        p = report_path + ext\n        if os.path.exists(p):\n            with open(p, encoding=\"utf-8\", errors=\"ignore\") as f:\n                return f.read()\n    if os.path.exists(report_path):\n        with open(report_path, encoding=\"utf-8\", errors=\"ignore\") as f:\n            return f.read()\n    return \"\"\n\n\ndef get_report_paths():\n    mapping = {}\n    for folder in (config.REPORTS_DIR, config.TEST_REPORTS_DIR):\n        if not os.path.exists(folder):\n            continue\n        for fn in os.listdir(folder):\n            uid = os.path.splitext(fn)[0]\n            mapping[uid] = os.path.join(folder, fn)\n    return mapping\n\n\n_REPORT_COL_HINTS = (\"report\", \"reports\", \"findings\", \"impression\", \"report_text\")\n\n\ndef find_report_column(df):\n    low = {c.lower(): c for c in df.columns}\n    for hint in _REPORT_COL_HINTS:\n        if hint in low:\n            return low[hint]\n    for c, orig in low.items():\n        if \"report\" in c:\n            return orig\n    return None\n\n\ndef reports_from_df(df):\n    col = find_report_column(df)\n    if col is None:\n        return {}\n    print(f\"[data] using report text from train.csv column '{col}'\")\n    uid_col = uid_column(df)\n    out = {}\n    for uid, txt in zip(df[uid_col], df[col]):\n        out[str(uid)] = txt if isinstance(txt, str) else \"\"\n    return out\n\n\ndef collect_reports(df, test_df=None):\n    \"\"\"Gather report text for train and (when available) test studies.\n\n    Returns (mapping, coverage) where coverage is the fraction of test studies\n    that actually resolved to non-empty text. A low value means the text branch\n    will be blank at inference time, which is what MODALITY_DROPOUT_TEXT guards\n    against -- so the caller reports it loudly rather than letting the fusion\n    head silently over-rely on a modality that disappears.\n    \"\"\"\n    mapping = {}\n    for uid, path in get_report_paths().items():\n        txt = load_report_text(path)\n        if txt:\n            mapping[str(uid)] = txt\n    if df is not None:\n        for uid, txt in reports_from_df(df).items():\n            if txt and uid not in mapping:\n                mapping[uid] = txt\n    if test_df is not None:\n        for uid, txt in reports_from_df(test_df).items():\n            if txt and uid not in mapping:\n                mapping[uid] = txt\n\n    coverage = 1.0\n    if test_df is not None and len(test_df):\n        col = uid_column(test_df)\n        hit = sum(1 for u in test_df[col] if mapping.get(str(u), \"\").strip())\n        coverage = hit / float(len(test_df))\n    return mapping, coverage\n\n\ndef _load_npy(path):\n    try:\n        return np.load(path, allow_pickle=False)\n    except Exception:\n        try:\n            return np.load(path, allow_pickle=True)\n        except Exception:\n            return None\n\n\ndef _load_npz(path):\n    try:\n        z = np.load(path, allow_pickle=False)\n        for k in z.files:\n            a = z[k]\n            if isinstance(a, np.ndarray) and a.size > 0:\n                return a\n        return None\n    except Exception:\n        return None\n\n\ndef _slice_sort_key(dcm):\n    try:\n        return (0, float(dcm.ImagePositionPatient[2]))\n    except Exception:\n        pass\n    try:\n        return (1, float(dcm.InstanceNumber))\n    except Exception:\n        return (2, 0.0)\n\n\ndef _series_label(dcm, fallback):\n    \"\"\"Best-effort sequence name, matched against config.SEQ_COLS when possible.\"\"\"\n    desc = \"\"\n    for attr in (\"SeriesDescription\", \"ProtocolName\", \"SequenceName\"):\n        v = getattr(dcm, attr, None)\n        if isinstance(v, str) and v.strip():\n            desc = v.strip()\n            break\n    if not desc:\n        return fallback\n    norm = re.sub(r\"[^a-z0-9]\", \"\", desc.lower())\n    for canonical in getattr(config, \"SEQ_COLS\", []):\n        if re.sub(r\"[^a-z0-9]\", \"\", canonical.lower()) in norm:\n            return canonical\n    return desc\n\n\ndef _read_dicom_series(file_paths):\n    \"\"\"Read one DICOM series and stack its slices into a (D, H, W) volume.\n\n    Slices within a series can differ spatially (640x640 vs 768x768), so every\n    plane is resampled onto the first valid slice's grid. Unreadable or\n    degenerate files are skipped rather than killing the batch.\"\"\"\n    items, label = [], None\n    for fp in file_paths:\n        try:\n            dcm = pydicom.dcmread(fp, force=True)\n        except Exception:\n            continue\n\n\n        try:\n            if not hasattr(dcm, \"pixel_array\"):\n                continue\n            arr = dcm.pixel_array.astype(np.float32)\n        except (ValueError, TypeError, RuntimeError, Exception) as e:\n            print(f\"[data] skipping bad slice {fp}: \"\n                  f\"{type(e).__name__}: {e}\", flush=True)\n            continue\n        if arr.ndim == 4 or (arr.ndim == 3 and arr.shape[-1] not in (1, 3, 4)):\n            continue\n        if arr.ndim == 3:\n            arr = arr.mean(axis=-1)\n        if arr.ndim != 2 or arr.size < 16:\n            continue\n        if label is None:\n            label = _series_label(dcm, None)\n        items.append((_slice_sort_key(dcm), np.ascontiguousarray(arr)))\n    if not items:\n        return None, None\n\n    items.sort(key=lambda t: t[0])\n    ref_h, ref_w = items[0][1].shape\n    if min(ref_h, ref_w) < 8 or max(ref_h, ref_w) > 1024:\n        ref_h, ref_w = 224, 224\n\n    out = []\n    for _, arr in items:\n        try:\n            out.append(arr if arr.shape == (ref_h, ref_w)\n                       else _resize2d(arr, ref_h, ref_w))\n        except Exception:\n            continue\n    if not out:\n        return None, None\n    return np.stack(out, axis=0), label\n\n\ndef _iter_dicom_groups(study_dir, max_depth=2):\n    groups = {}\n    base_depth = study_dir.rstrip(\"/\\\\\").count(os.sep)\n    for dirpath, dirnames, files in os.walk(study_dir):\n        depth = dirpath.rstrip(\"/\\\\\").count(os.sep) - base_depth\n        if depth >= max_depth:\n            dirnames[:] = []\n        dcm_files = [os.path.join(dirpath, f) for f in files\n                     if f.lower().endswith((\".dcm\", \".dicom\"))]\n        if dcm_files:\n            groups[dirpath] = dcm_files\n    return groups\n\n\ndef read_study_series(study_dir):\n    \"\"\"Return [(volume, series_label), ...] with series identity preserved.\n\n    The previous implementation concatenated every series into one axis-0 stack,\n    which destroyed sequence boundaries -- taking \"the middle 24 slices\" then\n    landed inside whichever series happened to sort into the centre, and which\n    series that was varied per study. Keeping series separate lets the sampler\n    spend its slice budget deliberately.\"\"\"\n    out = []\n    groups = _iter_dicom_groups(study_dir)\n    for i, d in enumerate(sorted(groups.keys())):\n        vol, label = _read_dicom_series(groups[d])\n        if vol is not None and vol.size > 0:\n            out.append((vol, label or f\"series{i}\"))\n    return out\n\n\ndef load_dicom_study(study_dir):\n    \"\"\"Backwards-compatible single-volume view of a study.\"\"\"\n    series = read_study_series(study_dir)\n    if not series:\n        return None\n    if len(series) == 1:\n        return series[0][0]\n    ref_h, ref_w = series[0][0].shape[1], series[0][0].shape[2]\n    if min(ref_h, ref_w) < 8 or max(ref_h, ref_w) > 4096:\n        ref_h, ref_w = 224, 224\n    matched = [resize_volume(v, ref_h, ref_w) for v, _ in series]\n    try:\n        return np.concatenate(matched, axis=0)\n    except Exception:\n        return matched[0]\n\n\ndef normalize_series(vol):\n    \"\"\"Per-series percentile normalisation to [0, 1].\n\n    Normalising over a whole multi-series study lets one bright sequence\n    compress every other sequence into a narrow band; each series gets its own\n    window instead.\"\"\"\n    vol = np.asarray(vol, dtype=np.float32)\n    finite = vol[np.isfinite(vol)]\n    if finite.size == 0:\n        return np.zeros_like(vol, dtype=np.float32)\n    lo, hi = np.percentile(finite, 0.5), np.percentile(finite, 99.5)\n    if hi - lo < 1e-6:\n        return np.zeros_like(vol, dtype=np.float32)\n    return np.clip((vol - lo) / (hi - lo), 0.0, 1.0).astype(np.float32)\n\n\ndef _normalize_vol(vol):\n    \"\"\"Kept for API compatibility with earlier revisions.\"\"\"\n    return normalize_series(vol)\n\n\ndef slice_saliency(vol):\n    \"\"\"Per-slice informativeness score in [0, 1].\n\n    Combines contrast (std) with tissue coverage (fraction of non-background\n    pixels). Empty end-slices, saturated localisers and padding score near\n    zero, so the sampler can spend its budget where anatomy actually is.\"\"\"\n    vol = np.asarray(vol, dtype=np.float32)\n    if vol.ndim != 3 or vol.shape[0] == 0:\n        return np.zeros((0,), dtype=np.float32)\n    flat = vol.reshape(vol.shape[0], -1)\n    std = flat.std(axis=1)\n    fg = (flat > 0.10).mean(axis=1)\n    score = std * np.sqrt(np.clip(fg, 1e-4, 1.0))\n    m = float(score.max())\n    if m < 1e-8:\n        return np.full(vol.shape[0], 1.0 / max(vol.shape[0], 1), dtype=np.float32)\n    return (score / m).astype(np.float32)\n\n\ndef _subsample_series(n, budget, min_keep):\n    \"\"\"Evenly spaced indices retaining at most `budget` of `n` slices.\"\"\"\n    keep = int(min(n, max(min_keep, budget)))\n    if keep >= n:\n        return np.arange(n)\n    return np.unique(np.linspace(0, n - 1, keep).astype(np.int64))\n\n\ndef _standardize_vol(vol, max_slices=None, max_hw=None):\n    \"\"\"Coerce an arbitrary array into a normalised (D, S, S) float32 volume.\"\"\"\n    if vol is None:\n        return None\n    max_slices = max_slices or getattr(config, \"MAX_SLICES\", 24)\n    max_slices = min(int(max_slices), 28)\n    max_hw = max_hw or getattr(config, \"VOLUME_CACHE_HW\",\n                               getattr(config, \"VOLUME_MAX_HW\", 256))\n    vol = np.asarray(vol)\n    if vol.ndim == 2:\n        vol = vol[None]\n    if vol.ndim == 4:\n        vol = vol[..., 0]\n    if vol.ndim != 3 or vol.shape[0] == 0:\n        return None\n    vol = vol.astype(np.float32, copy=False)\n    if vol.shape[0] > max_slices:\n        vol = vol[np.linspace(0, vol.shape[0] - 1, max_slices).astype(int)]\n    h, w = vol.shape[1], vol.shape[2]\n    if h != max_hw or w != max_hw:\n        vol = resize_volume(vol, max_hw, max_hw)\n    return normalize_series(vol)\n\n\ndef load_series_any(uid, volume_root, test_root=None):\n    \"\"\"Return [(volume, label), ...] for a study from any supported layout.\"\"\"\n    roots = [r for r in (volume_root, test_root) if r and os.path.isdir(r)]\n    for root in roots:\n        base = os.path.join(root, str(uid))\n        for p, loader in ((base + \".npy\", _load_npy), (base + \".npz\", _load_npz)):\n            if os.path.isfile(p):\n                a = loader(p)\n                if a is not None:\n                    return [(np.asarray(a, dtype=np.float32), \"npy\")]\n        if os.path.isfile(base):\n            a = _load_npy(base)\n            if a is not None:\n                return [(np.asarray(a, dtype=np.float32), \"npy\")]\n        if os.path.isdir(base):\n            series = read_study_series(base)\n            if series:\n                return series\n            parts = []\n            for p in sorted(glob.glob(os.path.join(base, \"*.npy\"))) + \\\n                    sorted(glob.glob(os.path.join(base, \"*.npz\"))):\n                a = _load_npz(p) if p.lower().endswith(\".npz\") else _load_npy(p)\n                if a is not None and np.asarray(a).size:\n                    parts.append((np.asarray(a, dtype=np.float32),\n                                  os.path.splitext(os.path.basename(p))[0]))\n            if parts:\n                return parts\n    return []\n\n\ndef load_volume_any(uid, volume_root, test_root=None):\n    \"\"\"Single standardised volume for a study, or None.\"\"\"\n    series = load_series_any(uid, volume_root, test_root)\n    if not series:\n        return None\n    vols = []\n    for v, _ in series:\n        s = _standardize_vol(v)\n        if s is not None:\n            vols.append(s)\n    if not vols:\n        return None\n    return vols[0] if len(vols) == 1 else np.concatenate(vols, axis=0)\n\n\ndef _load_from_root(uid, volume_root):\n    \"\"\"Retained for compatibility with earlier revisions.\"\"\"\n    return load_volume_any(uid, volume_root)\n\n\n_CACHE_INDEX_NAME = \"cache_index.json\"\n\n\ndef clear_volume_cache(cache_dir):\n    \"\"\"Remove a stale/corrupted cache directory before a fresh build.\n\n    Equivalent to `!rm -rf <cache_dir>`, but portable across OSes (the notebook\n    also runs on Windows locally). Dropping the whole folder avoids re-hitting a\n    corrupt `.npy`/`.meta.npz` left behind by an interrupted previous run.\"\"\"\n    if not cache_dir:\n        return\n    try:\n        import shutil\n        if os.path.isdir(cache_dir):\n            shutil.rmtree(cache_dir, ignore_errors=True)\n            print(f\"[cache] cleared stale cache: {cache_dir}\", flush=True)\n    except Exception:\n        pass\n\n\ndef clear_working_cache(working_dir=\"/kaggle/working\"):\n    \"\"\"Remove any volume-cache leftovers from the notebook Output directory.\n\n    Older revisions wrote the .npy/.meta.npz cache into /kaggle/working, which\n    bloats the Kaggle Output tab and can corrupt the exported submission zip.\n    The cache now lives in /tmp, so this one-time sweep at startup clears any\n    stale copy still sitting in the Output directory.\"\"\"\n    for name in (\"vol_cache\", \"vol_cache.npy\", \"vol_cache.meta.npz\"):\n        p = os.path.join(working_dir, name)\n        try:\n            import shutil\n            if os.path.isdir(p):\n                shutil.rmtree(p, ignore_errors=True)\n                print(f\"[cache] removed stale Output entry: {p}\", flush=True)\n            elif os.path.isfile(p):\n                os.remove(p)\n                print(f\"[cache] removed stale Output file: {p}\", flush=True)\n        except Exception:\n            pass\n\n\ndef cache_paths(uid, cache_dir):\n    base = os.path.join(cache_dir, str(uid))\n    return base + \".npy\", base + \".meta.npz\"\n\n\ndef is_cached(uid, cache_dir):\n    vol_p, meta_p = cache_paths(uid, cache_dir)\n    return os.path.isfile(vol_p) and os.path.isfile(meta_p)\n\n\ndef _to_storage_dtype(vol01, dtype_name):\n    if dtype_name == \"uint8\":\n        return np.clip(vol01 * 255.0 + 0.5, 0, 255).astype(np.uint8)\n    if dtype_name == \"float16\":\n        return vol01.astype(np.float16)\n    return vol01.astype(np.float32)\n\n\ndef from_storage_dtype(arr):\n    \"\"\"Inverse of _to_storage_dtype for an already-sliced array.\"\"\"\n    a = np.asarray(arr)\n    if a.dtype == np.uint8:\n        return a.astype(np.float32) * (1.0 / 255.0)\n    return a.astype(np.float32, copy=False)\n\n\ndef build_cache_entry(uid, volume_root, cache_dir, hw, max_slices,\n                      min_per_series, dtype_name):\n    \"\"\"Decode one study to (uid, n_slices, n_series) or (uid, 0, 0) on miss.\"\"\"\n    try:\n        if is_cached(uid, cache_dir):\n            vol_p, meta_p = cache_paths(uid, cache_dir)\n            try:\n                meta = np.load(meta_p)\n                return uid, int(meta[\"saliency\"].shape[0]), \\\n                    int(meta[\"offsets\"].shape[0]) - 1\n            except Exception:\n                pass\n\n        series = load_series_any(uid, volume_root)\n        if not series:\n            return uid, 0, 0\n\n        lengths = [max(int(np.asarray(v).shape[0]), 1) if np.asarray(v).ndim == 3\n                   else 1 for v, _ in series]\n        total = float(sum(lengths)) or 1.0\n\n        picked = []\n        for (vol, label), n in zip(series, lengths):\n            vol = np.asarray(vol, dtype=np.float32)\n            if vol.ndim == 2:\n                vol = vol[None]\n            if vol.ndim == 4:\n                vol = vol[..., 0]\n            if vol.ndim != 3 or vol.shape[0] == 0:\n                continue\n            budget = int(round(max_slices * (vol.shape[0] / total)))\n            idx = _subsample_series(vol.shape[0], budget, min_per_series)\n            vol = vol[idx]\n            vol = resize_volume(vol, hw, hw)\n            vol = normalize_series(vol)\n            picked.append((vol, str(label)))\n\n        if not picked:\n            return uid, 0, 0\n\n        total_kept = sum(int(v.shape[0]) for v, _ in picked)\n        if total_kept > max_slices:\n            n_new = [max(1, int(round(v.shape[0] * max_slices\n                                       / float(total_kept))))\n                     for v, _ in picked]\n            while sum(n_new) > max_slices:\n                j = int(np.argmax(n_new))\n                if n_new[j] <= 1:\n                    break\n                n_new[j] -= 1\n            hard_capped = []\n            for (v, lab), k in zip(picked, n_new):\n                if k < v.shape[0]:\n                    v = v[np.unique(np.linspace(0, v.shape[0] - 1,\n                                                k).astype(np.int64))]\n                hard_capped.append((v, lab))\n            picked = hard_capped\n\n        chunks, sal_chunks, offsets, labels = [], [], [0], []\n        for vol, label in picked:\n            chunks.append(_to_storage_dtype(vol, dtype_name))\n            sal_chunks.append(slice_saliency(vol))\n            offsets.append(offsets[-1] + vol.shape[0])\n            labels.append(label)\n\n        volume = np.concatenate(chunks, axis=0)\n        saliency = np.concatenate(sal_chunks, axis=0).astype(np.float32)\n        vol_p, meta_p = cache_paths(uid, cache_dir)\n        tmp_v = vol_p + \".tmp.npy\"\n        tmp_m = meta_p + \".tmp.npz\"\n        np.save(tmp_v, volume)\n        np.savez(tmp_m, saliency=saliency,\n                 offsets=np.asarray(offsets, dtype=np.int32),\n                 labels=np.asarray(labels))\n\n\n        try:\n            os.replace(tmp_v, vol_p)\n            os.replace(tmp_m, meta_p)\n        except FileNotFoundError:\n            pass\n        try:\n            for _p in (tmp_v, tmp_m):\n                if os.path.isfile(_p):\n                    os.remove(_p)\n        except OSError:\n            pass\n        return uid, int(volume.shape[0]), len(chunks)\n    except Exception as e:\n        print(f\"[cache] {uid}: {type(e).__name__}: {e}\", flush=True)\n        return uid, 0, 0\n\n\ndef _cache_worker(task):\n    return build_cache_entry(*task)\n\n\ndef _fit_cache_budget(n_studies, hw, max_slices, dtype_name, budget_gb):\n    \"\"\"Shrink resolution / slice count until the projected cache fits on disk.\"\"\"\n    bytes_per = {\"uint8\": 1, \"float16\": 2}.get(dtype_name, 4)\n\n    def gb(h, s):\n        return n_studies * s * h * h * bytes_per / (1024.0 ** 3)\n\n    if budget_gb <= 0 or gb(hw, max_slices) <= budget_gb:\n        return hw, max_slices\n    while max_slices > 24 and gb(hw, max_slices) > budget_gb:\n        max_slices -= 4\n    while hw > 160 and gb(hw, max_slices) > budget_gb:\n        hw -= 16\n    print(f\"[cache] projected size exceeded {budget_gb:.1f} GB budget; \"\n          f\"reduced to hw={hw}, slices={max_slices} \"\n          f\"(~{gb(hw, max_slices):.1f} GB)\")\n    return hw, max_slices\n\n\ndef build_volume_cache(uids, volume_root=None, cache_dir=None, workers=None,\n                       hw=None, max_slices=None, force=False, verbose=True):\n    \"\"\"Decode every study once into a uint8 mmap-friendly cache.\n\n    This is the single biggest throughput win in the pipeline: without it every\n    sample of every epoch re-reads and re-decodes a full DICOM study over\n    Kaggle's FUSE mount, then throws away all but ~24 slices.\"\"\"\n    volume_root = volume_root if volume_root is not None else config.VOLUME_ROOT\n    cache_dir = cache_dir or config.VOLUME_CACHE_DIR\n    if not volume_root or not os.path.isdir(volume_root):\n        if verbose:\n            print(f\"[cache] volume root unavailable ({volume_root}); \"\n                  f\"skipping cache build\")\n        return {\"cached\": 0, \"missing\": len(uids), \"dir\": cache_dir}\n\n    os.makedirs(cache_dir, exist_ok=True)\n    uids = [str(u) for u in dict.fromkeys(str(u) for u in uids)]\n    dtype_name = getattr(config, \"VOLUME_CACHE_DTYPE\", \"uint8\")\n    hw = hw or config.VOLUME_CACHE_HW\n    max_slices = min(int(max_slices or config.VOLUME_CACHE_MAX_SLICES),\n                     int(getattr(config, \"MAX_SLICES\", 24)))\n    hw, max_slices = _fit_cache_budget(\n        len(uids), hw, max_slices, dtype_name,\n        getattr(config, \"VOLUME_CACHE_BUDGET_GB\", 0.0))\n\n    _idx_p = os.path.join(cache_dir, _CACHE_INDEX_NAME)\n    if not force and os.path.isfile(_idx_p):\n        try:\n            with open(_idx_p, encoding=\"utf-8\") as f:\n                old = json.load(f)\n            if (old.get(\"max_slices\") != max_slices or old.get(\"hw\") != hw\n                    or old.get(\"dtype\") != dtype_name):\n                print(f\"[cache] existing cache built with \"\n                      f\"slices<={old.get('max_slices')} hw={old.get('hw')} \"\n                      f\"!= current slices<={max_slices} hw={hw}; \"\n                      f\"clearing to rebuild once\", flush=True)\n                clear_volume_cache(cache_dir)\n        except Exception:\n            pass\n\n    todo = uids if force else [u for u in uids if not is_cached(u, cache_dir)]\n    if verbose:\n        print(f\"[cache] dir={cache_dir}  hw={hw}  slices<={max_slices}  \"\n              f\"dtype={dtype_name}\")\n        print(f\"[cache] {len(uids) - len(todo)}/{len(uids)} already cached; \"\n              f\"decoding {len(todo)}\")\n\n    tasks = [(u, volume_root, cache_dir, hw, max_slices,\n              config.VOLUME_CACHE_MIN_PER_SERIES, dtype_name) for u in todo]\n    results = []\n    n_workers = workers if workers is not None else config.VOLUME_CACHE_WORKERS\n    n_workers = max(1, min(int(n_workers), os.cpu_count() or 1))\n\n    if tasks and n_workers > 1:\n        try:\n            import multiprocessing as mp\n\n            ctx = mp.get_context(\"spawn\" if os.name == \"nt\" else \"fork\")\n            with ctx.Pool(processes=n_workers) as pool:\n                for i, r in enumerate(pool.imap_unordered(_cache_worker, tasks,\n                                                          chunksize=4), 1):\n                    results.append(r)\n                    if verbose and (i % 200 == 0 or i == len(tasks)):\n                        print(f\"[cache] {i}/{len(tasks)} studies decoded\",\n                              flush=True)\n        except Exception as e:\n            print(f\"[cache] parallel prep unavailable ({e}); falling back to serial\")\n            results = []\n    if not results and tasks:\n        for i, t in enumerate(tasks, 1):\n            results.append(_cache_worker(t))\n            if verbose and (i % 100 == 0 or i == len(tasks)):\n                print(f\"[cache] {i}/{len(tasks)} studies decoded\", flush=True)\n\n    missing = [u for u, n, _ in results if n == 0]\n    index = {\"hw\": hw, \"max_slices\": max_slices, \"dtype\": dtype_name,\n             \"n_cached\": sum(1 for u in uids if is_cached(u, cache_dir)),\n             \"n_requested\": len(uids)}\n    try:\n        with open(os.path.join(cache_dir, _CACHE_INDEX_NAME), \"w\",\n                  encoding=\"utf-8\") as f:\n            json.dump(index, f)\n    except OSError:\n        pass\n\n    if verbose:\n        print(f\"[cache] ready: {index['n_cached']}/{len(uids)} studies\")\n        if missing:\n            print(f\"[cache] WARNING: {len(missing)} studies produced no volume \"\n                  f\"(e.g. {missing[:3]}). Those samples fall back to synthetic \"\n                  f\"noise and contribute nothing but label noise.\")\n    return {\"cached\": index[\"n_cached\"], \"missing\": len(missing),\n            \"dir\": cache_dir, \"hw\": hw, \"max_slices\": max_slices}\n\n\ndef load_cached(uid, cache_dir):\n    \"\"\"Memory-map a cached study. Returns (mmap_volume, saliency, offsets).\"\"\"\n    vol_p, meta_p = cache_paths(uid, cache_dir)\n    if not (os.path.isfile(vol_p) and os.path.isfile(meta_p)):\n        return None, None, None\n    try:\n        vol = np.load(vol_p, mmap_mode=\"r\")\n        meta = np.load(meta_p)\n        return vol, meta[\"saliency\"], meta[\"offsets\"]\n    except Exception:\n        return None, None, None\n\n\nif __name__ == \"__main__\":\n    import sys as _sys\n    import types as _types\n    _cell_module = _types.ModuleType(\"data_utils\")\n    _cell_module.__dict__.update(globals())\n    _sys.modules[\"data_utils\"] = _cell_module\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom collections import Counter\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\ntry:\n    import config\n    import data_utils as U\nexcept ModuleNotFoundError:\n    import sys as _sys\n    import types as _types\n    if \"config\" not in _sys.modules:\n        config = _types.ModuleType(\"config\")\n        config.__dict__.update(globals())\n        _sys.modules[\"config\"] = config\n    if \"data_utils\" not in _sys.modules:\n        U = _types.ModuleType(\"data_utils\")\n        U.__dict__.update(globals())\n        _sys.modules[\"data_utils\"] = U\n\ntry:\n    import cv2\n\n    cv2.setNumThreads(0)\n    _HAS_CV2 = True\nexcept Exception:  \n    _HAS_CV2 = False\n\n\n\n\n\n\n\n\n\n\ndef _tokenize(text):\n    return re.findall(r\"[a-z']+\", text.lower())\n\n\ndef _build_vocab(df, reports, min_df=None, vocab_size=None):\n    min_df = config.TEXT_MIN_DF if min_df is None else min_df\n    vocab_size = config.TEXT_VOCAB if vocab_size is None else vocab_size\n    uid_col = U.uid_column(df)\n    freq = Counter()\n    for uid in df[uid_col]:\n        freq.update(_tokenize(U.clean_text(reports.get(str(uid), \"\"))))\n    common = [w for w, c in freq.most_common(vocab_size - 2) if c >= min_df]\n    vocab = {\"<pad>\": 0, \"<unk>\": 1}\n    for i, w in enumerate(common, start=2):\n        vocab[w] = i\n    return vocab\n\n\ndef build_vocab(df, reports):\n    return _build_vocab(df, reports)\n\n\ndef _text_to_ids(text, vocab, max_len, word_dropout=0.0):\n    toks = _tokenize(text)[:max_len]\n    ids = [vocab.get(t, 1) for t in toks]\n    if word_dropout > 0.0 and ids:\n        ids = [1 if random.random() < word_dropout else t for t in ids]\n    if len(ids) < max_len:\n        ids += [0] * (max_len - len(ids))\n    return ids\n\n\n\n\n\ndef _allocate(total, weights, caps):\n    \"\"\"Split `total` picks across series by `weights`, respecting per-series caps.\n\n    Largest-remainder apportionment, then a spill loop so that a capped series\n    hands its surplus to the others instead of silently shrinking the batch.\"\"\"\n    k = len(weights)\n    if k == 0 or total <= 0:\n        return [0] * k\n    w = np.asarray(weights, dtype=np.float64)\n    caps = np.asarray(caps, dtype=np.int64)\n    if w.sum() <= 0:\n        w = np.ones(k, dtype=np.float64)\n    w = w / w.sum()\n\n    raw = w * total\n    alloc = np.floor(raw).astype(np.int64)\n    for i in np.argsort(-(raw - alloc)):\n        if alloc.sum() >= total:\n            break\n        alloc[i] += 1\n    alloc = np.minimum(alloc, caps)\n\n    \n    for _ in range(k + 1):\n        deficit = int(total - alloc.sum())\n        if deficit <= 0:\n            break\n        room = caps - alloc\n        if room.sum() <= 0:\n            break\n        for i in np.argsort(-room):\n            if deficit <= 0:\n                break\n            take = int(min(deficit, room[i]))\n            alloc[i] += take\n            deficit -= take\n    return alloc.tolist()\n\n\ndef _best_window(sal, want, jitter):\n    \"\"\"Start index of the highest-saliency contiguous window of length `want`.\n\n    Sliding-window sum via cumsum, so this is O(len) rather than O(len*want).\"\"\"\n    n = len(sal)\n    want = int(max(1, min(want, n)))\n    if want >= n:\n        return 0, n\n    csum = np.concatenate([[0.0], np.cumsum(sal, dtype=np.float64)])\n    sums = csum[want:] - csum[:-want]\n    start = int(np.argmax(sums))\n    if jitter > 0.0:\n        span = int(round(jitter * want))\n        if span > 0:\n            start = int(np.clip(start + random.randint(-span, span), 0, n - want))\n    return start, want\n\n\ndef select_slices(saliency, offsets, n_slices, train=True,\n                  strategy=None, jitter=None, temp=None):\n    \"\"\"Choose `n_slices` global slice indices from a multi-series study.\n\n    Budget is spread across series in proportion to their saliency mass, and\n    within each series the picks are drawn from that series' most informative\n    contiguous window. This replaces taking \"the middle N slices\" of a\n    concatenated volume, which sampled one arbitrary sequence per study.\"\"\"\n    strategy = strategy or getattr(config, \"SLICE_STRATEGY\", \"saliency\")\n    jitter = (config.SLICE_JITTER if jitter is None else jitter) if train else 0.0\n    temp = config.SALIENCY_TEMP if temp is None else temp\n\n    offsets = np.asarray(offsets, dtype=np.int64)\n    total = int(offsets[-1])\n    if total <= 0:\n        return np.zeros(0, dtype=np.int64)\n    if total <= n_slices:  \n        idx = np.arange(total, dtype=np.int64)\n        return idx[np.linspace(0, total - 1, n_slices).astype(np.int64)]\n\n    n_series = len(offsets) - 1\n    lengths = [int(offsets[k + 1] - offsets[k]) for k in range(n_series)]\n\n    if strategy == \"center\" or saliency is None:\n        mid = total // 2\n        half = n_slices // 2\n        start = int(np.clip(mid - half, 0, total - n_slices))\n        return np.arange(start, start + n_slices, dtype=np.int64)\n\n    saliency = np.asarray(saliency, dtype=np.float32)\n    mass = []\n    for k in range(n_series):\n        seg = saliency[offsets[k]:offsets[k + 1]]\n        mass.append(float(seg.sum()) if seg.size else 0.0)\n    \n    \n    weights = np.power(np.asarray(mass, dtype=np.float64) + 1e-6, temp)\n    budgets = _allocate(n_slices, weights, lengths)\n\n    picks = []\n    for k, b in enumerate(budgets):\n        if b <= 0:\n            continue\n        lo, hi = int(offsets[k]), int(offsets[k + 1])\n        seg = saliency[lo:hi]\n        \n        \n        want = int(min(hi - lo, max(b, round(b * 1.5))))\n        start, width = _best_window(seg, want, jitter)\n        local = np.linspace(start, start + width - 1, b).astype(np.int64)\n        picks.append(np.clip(local, 0, (hi - lo) - 1) + lo)\n\n    if not picks:\n        start = int(np.clip(total // 2 - n_slices // 2, 0, total - n_slices))\n        return np.arange(start, start + n_slices, dtype=np.int64)\n\n    idx = np.concatenate(picks)\n    if idx.shape[0] != n_slices:  \n        idx = idx[np.linspace(0, idx.shape[0] - 1, n_slices).astype(np.int64)]\n    return idx.astype(np.int64)\n\n\n\n\n\ndef _affine_stack(stack, angle, scale, tx, ty):\n    \"\"\"One shift-scale-rotate applied identically to every slice.\n\n    cv2 accepts up to 512 channels, so transposing (N,H,W) -> (H,W,N) turns N\n    separate warps into a single call -- roughly a 4x saving at N=24.\"\"\"\n    n, h, w = stack.shape\n    if _HAS_CV2 and n <= 512:\n        m = cv2.getRotationMatrix2D((w * 0.5, h * 0.5), angle, scale)\n        m[0, 2] += tx * w\n        m[1, 2] += ty * h\n        hwn = np.ascontiguousarray(stack.transpose(1, 2, 0))\n        out = cv2.warpAffine(hwn, m, (w, h), flags=cv2.INTER_LINEAR,\n                             borderMode=cv2.BORDER_REFLECT_101)\n        if out.ndim == 2:\n            out = out[:, :, None]\n        return np.ascontiguousarray(out.transpose(2, 0, 1), dtype=np.float32)\n\n    from skimage.transform import rotate as _sk_rotate\n    return np.stack([_sk_rotate(f, angle, order=1, mode=\"reflect\",\n                                preserve_range=True) for f in stack],\n                    axis=0).astype(np.float32)\n\n\ndef augment_stack(stack):\n    \"\"\"Volume-consistent augmentation: every slice gets the same geometry.\"\"\"\n    if random.random() < config.AUG_HFLIP_P:\n        stack = stack[:, :, ::-1]\n    if random.random() < config.AUG_SSR_P:\n        stack = _affine_stack(\n            np.ascontiguousarray(stack),\n            angle=random.uniform(-config.AUG_ROTATE, config.AUG_ROTATE),\n            scale=1.0 + random.uniform(-config.AUG_SCALE, config.AUG_SCALE),\n            tx=random.uniform(-config.AUG_SHIFT, config.AUG_SHIFT),\n            ty=random.uniform(-config.AUG_SHIFT, config.AUG_SHIFT))\n    if random.random() < config.AUG_BRIGHTNESS_P:\n        stack = stack * random.uniform(0.9, 1.1) + random.uniform(-0.05, 0.05)\n    if random.random() < config.AUG_GAMMA_P:\n        stack = np.power(np.clip(stack, 0.0, 1.0), random.uniform(0.85, 1.2))\n    if random.random() < config.AUG_NOISE_P:\n        stack = stack + np.random.normal(0.0, 0.02, stack.shape).astype(np.float32)\n    if random.random() < config.AUG_SLICE_DROPOUT_P and stack.shape[0] > 4:\n        k = random.randint(1, max(1, stack.shape[0] // 8))\n        stack = stack.copy()\n        stack[np.random.choice(stack.shape[0], k, replace=False)] = 0.0\n    return np.ascontiguousarray(np.clip(stack, 0.0, 1.0), dtype=np.float32)\n\n\n\n\n\nclass KneeDataset(Dataset):\n    \"\"\"Per-study knee volume + report text.\n\n    Volumes are read from the uint8 cache via mmap and fancy-indexed, so only\n    the N selected slices are ever paged in. Samples leave the worker as\n    float16 single-channel tensors; channel expansion happens on the GPU, which\n    cuts worker->main IPC traffic ~6x versus shipping float32 RGB.\"\"\"\n\n    def __init__(self, df, reports, vocab, mode=\"train\", volume_root=None,\n                 use_image=True, use_text=True, cache_dir=None):\n        self.df = df.reset_index(drop=True)\n        self.reports = reports or {}\n        self.vocab = vocab or {\"<pad>\": 0, \"<unk>\": 1}\n        self.mode = mode\n        self.volume_root = volume_root\n        self.use_image = use_image\n        self.use_text = use_text\n        self.targets = U.get_target_columns(df)\n        self.uid_col = U.uid_column(df)\n        self.n_seq = config.N_SEQ_IMAGES\n        self.img_size = config.IMG_SIZE\n        self.cache_dir = cache_dir if cache_dir is not None else (\n            config.VOLUME_CACHE_DIR if getattr(config, \"USE_VOLUME_CACHE\", True)\n            else None)\n        self._misses = 0\n        self._miss_warned = False\n        \n        \n        self._meta = {}\n\n    def __len__(self):\n        return len(self.df)\n\n    @property\n    def is_train(self):\n        return self.mode == \"train\"\n\n    \n    def _cached_entry(self, uid):\n        if self.cache_dir is None:\n            return None\n        entry = self._meta.get(uid)\n        if entry is not None:\n            return entry\n        vol, sal, offs = U.load_cached(uid, self.cache_dir)\n        if vol is None:\n            return None\n        entry = (vol, sal, offs)\n        \n        \n        if len(self._meta) < 512:\n            self._meta[uid] = entry\n        return entry\n\n    def _uncached_entry(self, uid):\n        try:\n            series = U.load_series_any(uid, self.volume_root)\n            if not series:\n                return None\n            chunks, sal, offs = [], [], [0]\n            for vol, _ in series:\n                v = U._standardize_vol(vol)\n                if v is None:\n                    continue\n                chunks.append(v)\n                sal.append(U.slice_saliency(v))\n                offs.append(offs[-1] + v.shape[0])\n            if not chunks:\n                return None\n            return (np.concatenate(chunks, axis=0),\n                    np.concatenate(sal).astype(np.float32),\n                    np.asarray(offs, dtype=np.int32))\n        except Exception as e:  \n            self._report_miss(uid)\n            return None\n\n    def _report_miss(self, uid):\n        self._misses += 1\n        if not self._miss_warned:\n            self._miss_warned = True\n            print(f\"[dataset] WARNING: no volume for study '{uid}' \"\n                  f\"(root={self.volume_root}, cache={self.cache_dir}). \"\n                  f\"Falling back to synthetic noise -- these samples add label \"\n                  f\"noise only. Run the cache prep step and check VOLUME_ROOT.\",\n                  flush=True)\n        elif self._misses in (10, 100, 1000):\n            print(f\"[dataset] {self._misses} volumes missing so far \"\n                  f\"({self.mode} split)\", flush=True)\n\n    def _fallback_stack(self):\n        \"\"\"A valid-shaped neutral stack so a failed study cannot crash the\n        DataLoader worker while still yielding a tensor with correct dims.\"\"\"\n        return np.full((self.n_seq, self.img_size, self.img_size), 0.0,\n                       dtype=np.float32)\n\n    def _load_stack(self, uid):\n        \"\"\"Return an (N, H, W) float32 stack in [0, 1] for one study.\"\"\"\n        try:\n            return self._load_stack_impl(uid)\n        except Exception as e:\n            \n            \n            self._report_miss(uid)\n            return self._fallback_stack()\n\n    def _load_stack_impl(self, uid):\n        entry = self._cached_entry(uid) or self._uncached_entry(uid)\n        if entry is None:\n            self._report_miss(uid)\n            return self._fallback_stack()\n\n        vol, sal, offs = entry\n        idx = select_slices(sal, offs, self.n_seq, train=self.is_train)\n        if idx.size == 0:\n            self._report_miss(uid)\n            return self._fallback_stack()\n\n        \n        \n        sel = np.asarray(vol[np.sort(idx)])\n        stack = U.from_storage_dtype(sel)\n        if stack.shape[1] != self.img_size or stack.shape[2] != self.img_size:\n            stack = U.resize_volume(stack, self.img_size, self.img_size)\n        return stack\n\n    \n    def _load_text(self, uid):\n        txt = U.clean_text(self.reports.get(uid, \"\"))\n        wd = config.TEXT_WORD_DROPOUT if self.is_train else 0.0\n        ids = _text_to_ids(txt, self.vocab, config.TEXT_MAX_LEN, word_dropout=wd)\n        text_ids = np.asarray(ids, dtype=np.int64)\n        return text_ids, (text_ids != 0).astype(np.float32)\n\n    def __getitem__(self, idx):\n        try:\n            return self._getitem_impl(idx)\n        except Exception as e:\n            \n            \n            print(f\"[dataset] WARNING: sample {idx} raised \"\n                  f\"{type(e).__name__}: {e}; returning fallback\", flush=True)\n            uid = str(self.df.iloc[idx][self.uid_col])\n            self._report_miss(uid)\n            image = None\n            if self.use_image:\n                stack = self._fallback_stack()\n                image = np.ascontiguousarray(stack[:, None, :, :],\n                                             dtype=np.float16)\n            text_ids = mask = None\n            if self.use_text:\n                text_ids, mask = self._load_text(uid)\n            if self.mode == \"test\":\n                return image, text_ids, mask, np.zeros(0, dtype=np.float32), uid\n            labels = np.zeros(config.NUM_TARGETS, dtype=np.float32)\n            if self.use_image:\n                labels = np.full(config.NUM_TARGETS, 0.5, dtype=np.float32)\n            return image, text_ids, mask, labels, uid\n\n    def _getitem_impl(self, idx):\n        row = self.df.iloc[idx]\n        uid = str(row[self.uid_col])\n\n        image = None\n        if self.use_image:\n            stack = self._load_stack(uid)\n            if float(stack.max() - stack.min()) < 1e-6:\n                stack = stack + np.random.normal(\n                    0.0, 0.01, stack.shape).astype(np.float32)\n            if self.is_train:\n                stack = augment_stack(stack)\n            \n            image = np.ascontiguousarray(stack[:, None, :, :], dtype=np.float16)\n\n        text_ids = mask = None\n        if self.use_text:\n            text_ids, mask = self._load_text(uid)\n\n        if self.mode == \"test\":\n            return image, text_ids, mask, np.zeros(0, dtype=np.float32), uid\n\n        labels = np.zeros(config.NUM_TARGETS, dtype=np.float32)\n        if self.targets:\n            labels = row[self.targets].values.astype(np.float32)\n        return image, text_ids, mask, labels, uid\n\n\nclass MultimodalCollator:\n    def __init__(self, use_image=True, use_text=True):\n        self.use_image = use_image\n        self.use_text = use_text\n\n    def __call__(self, batch):\n        images = texts = masks = labels = None\n        if self.use_image and batch[0][0] is not None:\n            \n            \n            images = torch.from_numpy(np.stack([b[0] for b in batch], axis=0))\n        if self.use_text and batch[0][1] is not None:\n            texts = torch.from_numpy(np.stack([b[1] for b in batch], axis=0))\n            masks = torch.from_numpy(np.stack([b[2] for b in batch], axis=0))\n        lab = batch[0][3]\n        if isinstance(lab, np.ndarray) and lab.size > 0:\n            labels = torch.from_numpy(np.stack([b[3] for b in batch], axis=0))\n        return {\"image\": images, \"text_ids\": texts, \"text_mask\": masks,\n                \"labels\": labels, \"uid\": [b[4] for b in batch]}\n\n\ndef make_loader(ds, batch_size=None, shuffle=False, drop_last=False):\n    \"\"\"RAM-safe DataLoader for Kaggle (small /dev/shm, tight memory limit).\n\n    Worker count is capped at 4 (default 2) so concurrent workers cannot push\n    system RAM / shared memory over the edge; persistent_workers is off so each\n    epoch's worker processes are torn down and their heap reclaimed; pin_memory\n    keeps the host->device copy fast when a GPU is present.\"\"\"\n    nw = min(int(getattr(config, \"NUM_WORKERS\", 2)), 4)\n    kwargs = dict(\n        batch_size=batch_size or config.BATCH_SIZE,\n        shuffle=shuffle,\n        drop_last=drop_last,\n        num_workers=nw,\n        collate_fn=MultimodalCollator(),\n        pin_memory=bool(getattr(config, \"PIN_MEMORY\", True)\n                        and config.DEVICE == \"cuda\"),\n    )\n    if nw > 0:\n        kwargs[\"prefetch_factor\"] = int(getattr(config, \"PREFETCH_FACTOR\", 2))\n        kwargs[\"persistent_workers\"] = bool(\n            getattr(config, \"PERSISTENT_WORKERS\", False))\n    return torch.utils.data.DataLoader(ds, **kwargs)\n\n\ndef configure_sharing_strategy():\n    \"\"\"Kaggle's /dev/shm is small; the default fd-based sharing strategy runs\n    out of descriptors on long runs and surfaces as a bus error or a silently\n    killed worker. Routing shared tensors through the filesystem avoids it.\"\"\"\n    want = getattr(config, \"SHARING_STRATEGY\", None)\n    if not want or os.name == \"nt\":\n        return\n    try:\n        import torch.multiprocessing as tmp\n\n        if want in tmp.get_all_sharing_strategies():\n            tmp.set_sharing_strategy(want)\n            print(f\"[loader] tensor sharing strategy = {want}\")\n    except Exception as e:\n        print(f\"[loader] could not set sharing strategy ({e})\")\n\n\nif __name__ == \"__main__\":\n    import sys as _sys\n    import types as _types\n    _cell_module = _types.ModuleType(\"dataset\")\n    _cell_module.__dict__.update(globals())\n    _sys.modules[\"dataset\"] = _cell_module\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n_TV_BLOCK_RE = re.compile(r\"^features\\.(\\d+)\\.(\\d+)\\.block\\.(.+)$\")\n\ntry:\n    import config\nexcept ModuleNotFoundError:\n    import sys as _sys\n    import types as _types\n    config = _types.ModuleType(\"config\")\n    config.__dict__.update(globals())\n    _sys.modules[\"config\"] = config\n\nWEIGHT_EXTS = (\".pth\", \".pt\", \".bin\", \".npz\")\n\n\ndef _norm(s):\n    return re.sub(r\"[^a-z0-9]\", \"\", str(s).lower())\n\n\ndef _search_roots():\n    roots = [\"/kaggle/input\",\n             os.path.join(os.environ.get(\"TORCH_HOME\", \"\"), \"hub\", \"checkpoints\"),\n             \"./pretrained\"]\n    return [r for r in roots if r and os.path.isdir(r)]\n\n\ndef _name_aliases(model_name):\n    \"\"\"Filename spellings a timm checkpoint for `model_name` might use.\"\"\"\n    t = _norm(model_name)\n    out = [t]\n    out.append(t.replace(\"efficientnet\", \"effnet\"))\n    out.append(t.replace(\"efficientnet\", \"tfefficientnet\"))\n    \n    m = re.match(r\"^(swin|convnext|efficientnet|resnet)([a-z]+)\", t)\n    if m:\n        out.append(m.group(1) + m.group(2))\n    return list(dict.fromkeys(a for a in out if a))\n\n\ndef find_pretrained_weights(model_name):\n    if config.PRETRAINED_WEIGHTS and os.path.isfile(config.PRETRAINED_WEIGHTS):\n        return config.PRETRAINED_WEIGHTS\n    aliases = _name_aliases(model_name)\n    cands = []\n    for root in _search_roots():\n        base_depth = root.rstrip(\"/\\\\\").count(os.sep)\n        for dirpath, dirnames, files in os.walk(root):\n            if dirpath.rstrip(\"/\\\\\").count(os.sep) - base_depth > 5:\n                dirnames[:] = []\n                continue\n            for fn in files:\n                if not fn.lower().endswith(WEIGHT_EXTS):\n                    continue\n                p = os.path.join(dirpath, fn)\n                try:\n                    if os.path.getsize(p) < 5 * 1024 * 1024:\n                        continue\n                except OSError:\n                    continue\n                cands.append(p)\n    \n    for alias in sorted(aliases, key=len, reverse=True):\n        for p in cands:\n            if alias in _norm(os.path.basename(p)):\n                return p\n    return None\n\n\n\n\n\ndef _reshape_to_ref(v, ref):\n    \"\"\"Adapt a checkpoint tensor onto the model's expected shape.\n\n    The only adaptation performed is input-channel folding on conv stems: a\n    pretrained RGB stem (O,3,k,k) is summed down to (O,1,k,k) so grayscale MRI\n    can reuse it. Summing (rather than slicing one channel) preserves the\n    response magnitude, which is the standard adaptation.\"\"\"\n    if tuple(v.shape) == tuple(ref.shape):\n        return v\n    if v.ndim == 4 and ref.ndim == 4 and v.shape[0] == ref.shape[0] \\\n            and v.shape[2:] == ref.shape[2:]:\n        cin_ckpt, cin_ref = v.shape[1], ref.shape[1]\n        if cin_ref == 1 and cin_ckpt > 1:\n            return v.sum(dim=1, keepdim=True)\n        if cin_ref > cin_ckpt and cin_ckpt == 1:\n            return v.repeat(1, cin_ref, 1, 1) / float(cin_ref)\n    return None\n\n\ndef load_state_dict_offline(model, path):\n    \"\"\"Load pretrained weights fully offline with automatic key-remapping.\"\"\"\n    if not path:\n        return False\n    if not os.path.exists(path):\n        print(f\"[offline] weights not found at {path}; using random init\")\n        return False\n    try:\n        raw = torch.load(path, map_location=\"cpu\")\n        if isinstance(raw, dict):\n            for key in (\"state_dict\", \"model\", \"model_state_dict\"):\n                if key in raw and isinstance(raw[key], dict):\n                    raw = raw[key]\n                    break\n        if not isinstance(raw, dict):\n            print(f\"[offline] unsupported weight format at {path}; random init\")\n            return False\n\n        msd = model.state_dict()\n        total = len(msd)\n        positions = _block_positions(model)\n\n        def lukemelas_fn(k):\n            return _lukemelas_candidates(k, positions)\n\n        strategies = (\n            (\"direct\", _direct_candidates),\n            (\"prefix-strip\", _prefix_candidates),\n            (\"lukemelas->timm\", lukemelas_fn),\n            (\"torchvision->timm\", _torchvision_candidates),\n        )\n\n        best_state, best_name, best_n, best_adapted = {}, \"none\", -1, 0\n        for name, cand_fn in strategies:\n            state, adapted = {}, 0\n            for k, v in raw.items():\n                if not hasattr(v, \"shape\"):\n                    continue\n                for nk in cand_fn(k):\n                    ref = msd.get(nk)\n                    if ref is None:\n                        continue\n                    fitted = _reshape_to_ref(v, ref)\n                    if fitted is None:\n                        continue\n                    if tuple(fitted.shape) != tuple(v.shape):\n                        adapted += 1\n                    state[nk] = fitted if fitted.dtype == ref.dtype \\\n                        else fitted.to(ref.dtype)\n                    break\n            if len(state) > best_n:\n                best_state, best_name, best_n, best_adapted = \\\n                    state, name, len(state), adapted\n            if best_n >= 0.99 * total:\n                break\n\n        missing, unexpected = model.load_state_dict(best_state, strict=False)\n        ok = best_n >= 0.5 * total\n        if ok:\n            extra = (f\", in_chans-adapted={best_adapted}\" if best_adapted else \"\")\n            print(f\"[offline] pretrained weights loaded from {os.path.basename(path)} \"\n                  f\"via '{best_name}' mapping (matched={best_n}/{total}, \"\n                  f\"missing={len(missing)}, unexpected={len(unexpected)}{extra})\")\n        else:\n            sample = list(raw.keys())[:6]\n            print(f\"[offline] weights at {os.path.basename(path)} do not match \"\n                  f\"architecture (best '{best_name}' matched={best_n}/{total}); \"\n                  f\"random init. Sample checkpoint keys: {sample}\")\n        return ok\n    except Exception as e:\n        print(f\"[offline] failed to load pretrained weights ({e}); random init\")\n        return False\n\n\ndef _direct_candidates(k):\n    return [k]\n\n\ndef _prefix_candidates(k):\n    cands = [k]\n    for pfx in (\"module.\", \"model.\", \"backbone.\", \"encoder.\", \"net.\"):\n        if k.startswith(pfx):\n            cands.append(k[len(pfx):])\n    return cands\n\n\n_BLOCK_ALIASES = (\n    (\"_depthwise_conv\", \"conv_dw\"),\n    (\"_expand_conv\", \"conv_pw\"),\n    (\"_project_conv\", \"conv_pwl\"),\n    (\"_bn0\", \"bn1\"),\n    (\"_bn1\", \"bn2\"),\n    (\"_bn2\", \"bn3\"),\n    (\"_se_reduce\", \"se.conv_reduce\"),\n    (\"_se_expand\", \"se.conv_expand\"),\n)\n\n\ndef _lukemelas_candidates(k, block_positions=None):\n    out = list(_prefix_candidates(k))\n    c = out[-1]\n    if c.startswith(\"_conv_stem\"):\n        out.append(\"conv_stem\" + c[len(\"_conv_stem\"):])\n    elif c.startswith(\"_conv_head\"):\n        out.append(\"conv_head\" + c[len(\"_conv_head\"):])\n    elif c.startswith(\"_bn0\"):\n        out.append(\"bn1\" + c[len(\"_bn0\"):])\n    elif c.startswith(\"_bn1\"):\n        out.append(\"bn2\" + c[len(\"_bn1\"):])\n    elif c.startswith(\"_fc.\"):\n        out.append(\"fc.\" + c[4:])\n    elif c.startswith(\"_blocks.\"):\n        parts = c.split(\".\")\n        idx, suffix = int(parts[1]), \".\".join(parts[2:])\n        mapped = None\n        for old, new in _BLOCK_ALIASES:\n            if suffix.startswith(old):\n                mapped = new + suffix[len(old):]\n                break\n        if mapped is not None:\n            if block_positions is not None and idx < len(block_positions):\n                stage, blk = block_positions[idx]\n                out.append(f\"blocks.{stage}.{blk}.{mapped}\")\n            out.append(f\"blocks.{idx}.{mapped}\")\n    return out\n\n\ndef _block_positions(model):\n    positions = []\n    for key in model.state_dict().keys():\n        parts = key.split(\".\")\n        if (len(parts) >= 4 and parts[0] == \"blocks\"\n                and parts[1].isdigit() and parts[2].isdigit()):\n            pos = (int(parts[1]), int(parts[2]))\n            if pos not in positions:\n                positions.append(pos)\n    return positions\n\n\ndef _torchvision_candidates(k):\n    out = list(_prefix_candidates(k))\n    c = out[-1]\n    m = _TV_BLOCK_RE.match(c)\n    if m:\n        stage, blk, rest = int(m.group(1)), int(m.group(2)), m.group(3)\n        if 1 <= stage <= 7:\n            base = f\"blocks.{stage - 1}.{blk}.\"\n            for a, b in ((\"0.0.\", \"conv_pw.\"), (\"0.1.\", \"bn1.\"),\n                         (\"1.0.\", \"conv_dw.\"), (\"1.1.\", \"bn2.\"),\n                         (\"2.fc1.\", \"se.conv_reduce.\"),\n                         (\"2.fc2.\", \"se.conv_expand.\"),\n                         (\"3.0.\", \"conv_pwl.\"), (\"3.1.\", \"bn3.\")):\n                if rest.startswith(a):\n                    out.append(base + b + rest[len(a):])\n                    break\n        return out\n    if c.startswith(\"features.0.0.\"):\n        out.append(\"conv_stem.\" + c[len(\"features.0.0.\"):])\n    elif c.startswith(\"features.0.1.\"):\n        out.append(\"bn1.\" + c[len(\"features.0.1.\"):])\n    elif c.startswith(\"features.8.0.\"):\n        out.append(\"conv_head.\" + c[len(\"features.8.0.\"):])\n    elif c.startswith(\"features.8.1.\"):\n        out.append(\"bn2.\" + c[len(\"features.8.1.\"):])\n    return out\n\n\n\n\n\n_TRANSFORMER_HINTS = (\"swin\", \"vit\", \"deit\", \"beit\", \"coat\", \"maxvit\")\n\n\ndef is_transformer_backbone(name=None):\n    n = (name or config.MODEL_NAME).lower()\n    return any(h in n for h in _TRANSFORMER_HINTS)\n\n\ndef _make_backbone():\n    \"\"\"Create a timm backbone with pooled features and no classifier.\"\"\"\n    import timm\n\n    name = config.MODEL_NAME\n    in_chans = int(getattr(config, \"IN_CHANS\", 1))\n    path = find_pretrained_weights(name)\n    if path is None:\n        alt = find_pretrained_weights(\"resnet50\")\n        if alt is not None:\n            print(f\"[offline] no weights found for '{name}'; falling back to \"\n                  f\"resnet50 weights ({os.path.basename(alt)})\")\n            name, path = \"resnet50\", alt\n\n    try:\n        model = timm.create_model(name, pretrained=False, num_classes=0,\n                                  in_chans=in_chans)\n    except Exception as e:\n        print(f\"[model] timm could not build '{name}' ({e}); \"\n              f\"falling back to efficientnet_b0\")\n        name = \"efficientnet_b0\"\n        model = timm.create_model(name, pretrained=False, num_classes=0,\n                                  in_chans=in_chans)\n        path = find_pretrained_weights(name)\n\n    if path is not None:\n        load_state_dict_offline(model, path)\n    else:\n        print(f\"[offline] WARNING: no pretrained weights found for '{name}'. \"\n              f\"Vision starts RANDOM -> weak score. Attach a Kaggle dataset \"\n              f\"containing a timm '{name}' .pth file (see beginner_user_manual.md).\")\n\n    feat = int(getattr(model, \"num_features\", config.ENCODER_FEAT))\n    if getattr(config, \"GRAD_CHECKPOINT\", False) and \\\n            hasattr(model, \"set_grad_checkpointing\"):\n        try:\n            model.set_grad_checkpointing(enable=True)\n            print(f\"[model] gradient checkpointing enabled on {name}\")\n        except Exception:\n            pass\n    return model, feat, name\n\n\n\n\n\nclass SliceTransformer(nn.Module):\n    \"\"\"Self-attention across the slice axis.\n\n    A per-slice CNN sees each plane independently; this lets the N slice\n    embeddings exchange information so the head can reason about extent along\n    the depth axis (a meniscal tear spanning several slices vs. one artefact).\"\"\"\n\n    def __init__(self, dim, layers=2, heads=8, dropout=0.1, max_slices=96):\n        super().__init__()\n        self.pos = nn.Parameter(torch.zeros(1, max_slices, dim))\n        nn.init.trunc_normal_(self.pos, std=0.02)\n        layer = nn.TransformerEncoderLayer(\n            d_model=dim, nhead=heads, dim_feedforward=dim * 2,\n            dropout=dropout, activation=\"gelu\", batch_first=True,\n            norm_first=True)\n        self.enc = nn.TransformerEncoder(layer, num_layers=layers)\n        self.norm = nn.LayerNorm(dim)\n\n    def forward(self, x):\n        n = x.size(1)\n        if n > self.pos.size(1):  \n            pos = F.interpolate(self.pos.transpose(1, 2), size=n,\n                                mode=\"linear\", align_corners=False).transpose(1, 2)\n        else:\n            pos = self.pos[:, :n]\n        return self.norm(self.enc(x + pos))\n\n\nclass MultiPool3D(nn.Module):\n    \"\"\"Adaptive attention + average + max pooling over the slice axis.\n\n    Average pooling captures diffuse findings (effusion, synovitis), max\n    pooling captures focal ones (fracture, contusion), and the learned\n    attention weights slices by relevance. Concatenating all three and\n    projecting back down consistently beats any single pooling operator.\"\"\"\n\n    def __init__(self, dim, hidden=256, out_dim=None, dropout=0.1):\n        super().__init__()\n        out_dim = out_dim or dim\n        self.score = nn.Sequential(\n            nn.Linear(dim, hidden), nn.Tanh(), nn.Linear(hidden, 1))\n        self.proj = nn.Sequential(\n            nn.LayerNorm(dim * 3),\n            nn.Dropout(dropout),\n            nn.Linear(dim * 3, out_dim),\n        )\n\n    def forward(self, x):\n        w = torch.softmax(self.score(x).squeeze(-1), dim=1)   \n        attn = torch.einsum(\"bn,bnd->bd\", w, x)\n        avg = x.mean(dim=1)\n        mx = x.amax(dim=1)\n        return self.proj(torch.cat([attn, avg, mx], dim=1)), w\n\n\nclass SliceAttention(nn.Module):\n    \"\"\"Retained for checkpoint/API compatibility with earlier revisions.\"\"\"\n\n    def __init__(self, dim, hidden=256):\n        super().__init__()\n        self.proj = nn.Linear(dim, hidden)\n        self.query = nn.Parameter(torch.randn(hidden) * 0.02)\n        self.v = nn.Linear(hidden, 1)\n\n    def forward(self, x):\n        h = torch.tanh(self.proj(x))\n        scores = self.v(h * self.query).squeeze(-1)\n        w = torch.softmax(scores, dim=1)\n        return (x * w.unsqueeze(-1)).sum(dim=1)\n\n\n\n\n\nclass TextEncoder(nn.Module):\n    \"\"\"BiGRU over report tokens, returning both token states and a pooled vector.\"\"\"\n\n    def __init__(self, vocab_size=None, hidden=None, dropout=None):\n        super().__init__()\n        vocab_size = vocab_size or config.TEXT_VOCAB\n        hidden = hidden or config.TEXT_HIDDEN\n        dropout = config.DROPOUT if dropout is None else dropout\n        self.hidden = hidden\n        self.tok_embed = nn.Embedding(vocab_size, hidden // 2, padding_idx=0)\n        self.enc = nn.GRU(hidden // 2, hidden // 2, num_layers=1,\n                          batch_first=True, bidirectional=True)\n        self.norm = nn.LayerNorm(hidden)\n        self.proj = nn.Linear(hidden, hidden)\n        self.drop = nn.Dropout(dropout)\n\n    def forward(self, ids, mask):\n        out, _ = self.enc(self.tok_embed(ids))\n        tokens = self.norm(out)\n        if mask is not None:\n            m = mask.unsqueeze(-1).to(tokens.dtype)\n            pooled = (tokens * m).sum(dim=1) / m.sum(dim=1).clamp(min=1.0)\n        else:\n            pooled = tokens.mean(dim=1)\n        return tokens, self.proj(self.drop(pooled))\n\n\n\n\n\nclass GatedCrossAttentionFusion(nn.Module):\n    \"\"\"Image-queried cross-attention over report tokens, then a learned gate.\n\n    The image embedding queries the report so the model can align \"lateral\n    meniscus\" in the text with the slices that show it, instead of blending two\n    pooled vectors blindly. The sigmoid gate then decides per-feature how much\n    to trust each modality, which is what lets the network degrade gracefully\n    when a report is missing rather than emitting garbage.\"\"\"\n\n    def __init__(self, img_dim, txt_dim, dim=None, heads=None, dropout=None):\n        super().__init__()\n        dim = dim or config.FUSION_DIM\n        heads = heads or config.FUSION_HEADS\n        dropout = config.DROPOUT if dropout is None else dropout\n        while dim % heads != 0 and heads > 1:\n            heads -= 1\n\n        self.dim = dim\n        self.img_proj = nn.Sequential(nn.Linear(img_dim, dim), nn.LayerNorm(dim))\n        self.txt_proj = nn.Sequential(nn.Linear(txt_dim, dim), nn.LayerNorm(dim))\n        self.tok_proj = nn.Sequential(nn.Linear(txt_dim, dim), nn.LayerNorm(dim))\n        self.attn = nn.MultiheadAttention(dim, heads, dropout=dropout,\n                                          batch_first=True)\n        self.attn_norm = nn.LayerNorm(dim)\n        self.gate = nn.Sequential(\n            nn.Linear(dim * 2, dim), nn.LayerNorm(dim), nn.Sigmoid())\n        self.drop = nn.Dropout(dropout)\n        self.out_dim = dim * 3\n\n    def forward(self, img_feat, txt_tokens, txt_pooled, txt_mask, txt_keep):\n        img = self.img_proj(img_feat)                               \n\n        if txt_tokens is None:\n            txt = torch.zeros_like(img)\n        else:\n            txt = self.txt_proj(txt_pooled)\n            kv = self.tok_proj(txt_tokens)                           \n            pad = None\n            if txt_mask is not None:\n                pad = txt_mask <= 0.5\n                \n                \n                \n                \n                pad = pad.clone()\n                pad[:, 0] = False\n            att, _ = self.attn(img.unsqueeze(1), kv, kv, key_padding_mask=pad,\n                               need_weights=False)\n            txt = self.attn_norm(txt + att.squeeze(1))\n\n        if txt_keep is not None:\n            txt = txt * txt_keep.view(-1, 1).to(txt.dtype)\n\n        g = self.gate(torch.cat([img, txt], dim=1))\n        fused = g * img + (1.0 - g) * txt\n        return self.drop(torch.cat([img, txt, fused], dim=1))\n\n\n\n\n\nclass RSNAModel(nn.Module):\n    def __init__(self, vocab_size=None, num_targets=None):\n        super().__init__()\n        num_targets = num_targets or config.NUM_TARGETS\n        self.backbone, feat, self.backbone_name = _make_backbone()\n        self.feat_dim = feat\n        self.in_chans = int(getattr(config, \"IN_CHANS\", 1))\n        self.is_transformer = is_transformer_backbone(self.backbone_name)\n        self.use_channels_last = bool(getattr(config, \"CHANNELS_LAST\", True)\n                                      and not self.is_transformer)\n\n        slice_dim = min(feat, config.FUSION_HIDDEN)\n        self.slice_dim = slice_dim\n        self.slice_in = (nn.Identity() if slice_dim == feat\n                         else nn.Linear(feat, slice_dim))\n        self.slice_ctx = SliceTransformer(\n            slice_dim,\n            layers=config.SLICE_TRANSFORMER_LAYERS,\n            heads=config.SLICE_TRANSFORMER_HEADS,\n            dropout=min(config.DROPOUT, 0.2)) \\\n            if config.SLICE_TRANSFORMER_LAYERS > 0 else None\n        self.pool = MultiPool3D(slice_dim, out_dim=slice_dim,\n                                dropout=min(config.DROPOUT, 0.2))\n\n        self.text_encoder = TextEncoder(vocab_size=vocab_size or config.TEXT_VOCAB)\n        self.fusion = GatedCrossAttentionFusion(slice_dim, config.TEXT_HIDDEN)\n\n        self.head = nn.Sequential(\n            nn.LayerNorm(self.fusion.out_dim),\n            nn.Dropout(config.DROPOUT),\n            nn.Linear(self.fusion.out_dim, config.FUSION_HIDDEN),\n            nn.GELU(),\n            nn.Dropout(config.DROPOUT),\n            nn.Linear(config.FUSION_HIDDEN, num_targets),\n        )\n        \n        \n        \n        nn.init.constant_(self.head[-1].bias, -2.0)\n\n    \n    def _to_backbone_input(self, images):\n        \"\"\"(B, N, C, H, W) -> (B*N, in_chans, H, W) laid out for the GPU.\"\"\"\n        b, n = images.shape[:2]\n        x = images.reshape(b * n, *images.shape[2:])\n        if x.size(1) == 1 and self.in_chans > 1:\n            x = x.expand(-1, self.in_chans, -1, -1)\n        elif x.size(1) > self.in_chans:\n            x = x[:, :self.in_chans]\n        if self.use_channels_last:\n            x = x.contiguous(memory_format=torch.channels_last)\n        else:\n            x = x.contiguous()\n        return x, b, n\n\n    def image_embedding(self, images):\n        x, b, n = self._to_backbone_input(images)\n        feats = self.backbone(x)                 \n        if feats.ndim > 2:                       \n            feats = feats.flatten(2).mean(-1)\n        feats = self.slice_in(feats).view(b, n, -1)\n        if self.slice_ctx is not None:\n            feats = self.slice_ctx(feats)\n        pooled, _ = self.pool(feats)\n        return pooled\n\n    \n    def _modality_keep(self, batch, text_mask, device, dtype):\n        \"\"\"Per-sample keep mask for the text branch.\n\n        Combines genuine absence (an all-pad report) with training-time modality\n        dropout. Reports in this competition restate the findings almost\n        verbatim, so without this the fusion head learns to read the answer off\n        the text and then collapses on any study whose report is unavailable.\"\"\"\n        keep = torch.ones(batch, device=device, dtype=dtype)\n        if text_mask is not None:\n            keep = (text_mask.sum(dim=1) > 0).to(dtype)\n        if self.training and config.MODALITY_DROPOUT_TEXT > 0:\n            drop = torch.rand(batch, device=device) < config.MODALITY_DROPOUT_TEXT\n            keep = keep * (~drop).to(dtype)\n        return keep\n\n    def forward(self, images=None, text_ids=None, text_mask=None):\n        if images is None and text_ids is None:\n            raise ValueError(\"RSNAModel.forward needs images, text_ids, or both\")\n\n        img_feat = None\n        if images is not None:\n            img_feat = self.image_embedding(images)\n            if self.training and config.MODALITY_DROPOUT_IMAGE > 0:\n                b = img_feat.size(0)\n                drop = torch.rand(b, device=img_feat.device) \\\n                    < config.MODALITY_DROPOUT_IMAGE\n                img_feat = img_feat * (~drop).view(-1, 1).to(img_feat.dtype)\n\n        if text_ids is not None:\n            tokens, pooled = self.text_encoder(text_ids, text_mask)\n            ref = img_feat if img_feat is not None else pooled\n            keep = self._modality_keep(ref.size(0), text_mask, ref.device,\n                                       ref.dtype)\n        else:\n            tokens = pooled = None\n            keep = None\n\n        if img_feat is None:\n            \n            \n            img_feat = torch.zeros(pooled.size(0), self.slice_dim,\n                                   device=pooled.device, dtype=pooled.dtype)\n        fused = self.fusion(img_feat, tokens, pooled, text_mask, keep)\n        return self.head(fused)\n\n\n\n\n\nclass AsymmetricLoss(nn.Module):\n    \"\"\"Asymmetric loss for multi-label imbalance (Ben-Baruch et al., 2020).\n\n    Down-weights easy negatives with gamma_neg while leaving positives nearly\n    unfocused, and hard-clips very confident negatives out of the gradient\n    entirely. For 12 targets whose prevalence spans roughly 1%-40% this is\n    consistently stronger than BCE with pos_weight, which has to trade\n    stability against how far it can push the rare classes.\"\"\"\n\n    def __init__(self, gamma_neg=None, gamma_pos=None, clip=None, eps=1e-8):\n        super().__init__()\n        self.gamma_neg = config.ASL_GAMMA_NEG if gamma_neg is None else gamma_neg\n        self.gamma_pos = config.ASL_GAMMA_POS if gamma_pos is None else gamma_pos\n        self.clip = config.ASL_CLIP if clip is None else clip\n        self.eps = eps\n\n    def forward(self, logits, targets):\n        logits = logits.float()\n        targets = targets.float()\n        xs_pos = torch.sigmoid(logits)\n        xs_neg = 1.0 - xs_pos\n        if self.clip and self.clip > 0:\n            xs_neg = (xs_neg + self.clip).clamp(max=1.0)\n\n        loss = targets * torch.log(xs_pos.clamp(min=self.eps)) \\\n            + (1.0 - targets) * torch.log(xs_neg.clamp(min=self.eps))\n\n        if self.gamma_neg > 0 or self.gamma_pos > 0:\n            \n            \n            \n            with torch.no_grad():\n                pt = xs_pos * targets + xs_neg * (1.0 - targets)\n                gamma = self.gamma_pos * targets \\\n                    + self.gamma_neg * (1.0 - targets)\n                w = torch.pow(1.0 - pt, gamma)\n            loss = loss * w\n        return -loss.mean()\n\n\nclass FocalBCELoss(nn.Module):\n    def __init__(self, gamma=None, pos_weight=None):\n        super().__init__()\n        self.gamma = config.FOCAL_GAMMA if gamma is None else gamma\n        self.register_buffer(\n            \"pos_weight\",\n            None if pos_weight is None else pos_weight.float(),\n            persistent=False)\n\n    def forward(self, logits, targets):\n        logits = logits.float()\n        targets = targets.float()\n        bce = F.binary_cross_entropy_with_logits(\n            logits, targets, reduction=\"none\",\n            pos_weight=self.pos_weight)\n        with torch.no_grad():\n            p = torch.sigmoid(logits)\n            pt = p * targets + (1.0 - p) * (1.0 - targets)\n            w = torch.pow(1.0 - pt, self.gamma)\n        return (bce * w).mean()\n\n\nclass StableBCEWithLogits(nn.Module):\n    \"\"\"BCEWithLogitsLoss that upcasts before the loss, for AMP correctness.\"\"\"\n\n    def __init__(self, pos_weight=None):\n        super().__init__()\n        self.register_buffer(\n            \"pos_weight\",\n            None if pos_weight is None else pos_weight.float(),\n            persistent=False)\n\n    def forward(self, logits, targets):\n        return F.binary_cross_entropy_with_logits(\n            logits.float(), targets.float(), pos_weight=self.pos_weight)\n\n\ndef compute_pos_weight(targets, cap=None):\n    \"\"\"neg/pos ratio per target, clipped so a 1%-prevalence class cannot\n    dominate the gradient.\"\"\"\n    import numpy as np\n\n    cap = config.POS_WEIGHT_MAX if cap is None else cap\n    pos = np.asarray(targets, dtype=np.float64).sum(axis=0)\n    neg = np.asarray(targets).shape[0] - pos\n    w = np.clip(neg / np.maximum(pos, 1.0), 1.0, cap)\n    return torch.tensor(w, dtype=torch.float32)\n\n\ndef build_criterion(targets=None, kind=None, device=None):\n    \"\"\"Instantiate the configured loss, wiring pos_weight from label prevalence.\"\"\"\n    kind = (kind or getattr(config, \"LOSS\", \"bce\")).lower()\n    device = device or config.DEVICE\n    pw = None\n    if targets is not None and len(targets):\n        pw = compute_pos_weight(targets).to(device)\n\n    if kind == \"asl\":\n        crit = AsymmetricLoss()\n    elif kind == \"focal\":\n        crit = FocalBCELoss(pos_weight=pw)\n    else:\n        crit = StableBCEWithLogits(pos_weight=pw)\n    print(f\"[loss] {kind}\" + (f\" (pos_weight max={float(pw.max()):.2f})\"\n                              if pw is not None and kind != \"asl\" else \"\"))\n    return crit.to(device)\n\n\n\n\n\ndef canonical_key(key):\n    \"\"\"Drop every DataParallel 'module' path segment, wherever it sits.\n\n    Because DataParallel is applied to the backbone rather than the whole\n    model, wrapped keys look like 'backbone.module.conv_stem.weight'. Saving\n    canonical names keeps a checkpoint interchangeable between a single-GPU and\n    a dual-T4 run.\"\"\"\n    return \".\".join(p for p in key.split(\".\") if p != \"module\")\n\n\ndef canonical_state_dict(model):\n    return {canonical_key(k): v for k, v in model.state_dict().items()}\n\n\ndef load_canonical_state_dict(model, state, strict=False, verbose=True):\n    lookup = {canonical_key(k): k for k in model.state_dict().keys()}\n    own = model.state_dict()\n    remapped, skipped = {}, []\n    for k, v in state.items():\n        target = lookup.get(canonical_key(k))\n        if target is None:\n            skipped.append(k)\n            continue\n        if hasattr(v, \"shape\") and tuple(own[target].shape) != tuple(v.shape):\n            skipped.append(k)\n            continue\n        remapped[target] = v\n    missing, unexpected = model.load_state_dict(remapped, strict=strict)\n    if verbose and (skipped or missing):\n        print(f\"[ckpt] loaded {len(remapped)}/{len(own)} tensors \"\n              f\"(skipped={len(skipped)}, missing={len(missing)}, \"\n              f\"unexpected={len(unexpected)})\")\n    return missing, unexpected\n\n\n\n\n\ndef enable_data_parallel(model, verbose=True):\n    \"\"\"Shard the vision backbone across every visible GPU.\n\n    Only the backbone is wrapped, deliberately. The per-slice tensor is\n    (B*N, C, H, W) -- the largest tensor in the graph and perfectly splittable\n    on dim 0 -- while the GRU, cross-attention and head are small and stay on\n    cuda:0. Wrapping the whole model instead would replicate the GRU on every\n    step (slow, and it triggers the non-contiguous-weights warning) and would\n    scatter/gather several small tensors per forward for no benefit.\"\"\"\n    if not (getattr(config, \"DATA_PARALLEL\", True)\n            and config.DEVICE == \"cuda\" and torch.cuda.device_count() > 1):\n        return model\n    target = model.module if isinstance(model, nn.DataParallel) else model\n    if isinstance(getattr(target, \"backbone\", None), nn.DataParallel):\n        return model\n    n = torch.cuda.device_count()\n    target.backbone = nn.DataParallel(target.backbone,\n                                      device_ids=list(range(n)))\n    if verbose:\n        eff = config.BATCH_SIZE * config.N_SEQ_IMAGES\n        print(f\"[gpu] DataParallel over {n} devices on the vision backbone \"\n              f\"({eff} slices/step -> ~{eff // n} per GPU)\")\n    return model\n\n\ndef build_model(vocab_size=None, num_targets=None, parallel=True):\n    model = RSNAModel(vocab_size=vocab_size, num_targets=num_targets)\n    model = model.to(config.DEVICE)\n    if getattr(config, \"CHANNELS_LAST\", True) and config.DEVICE == \"cuda\" \\\n            and not model.is_transformer:\n        model = model.to(memory_format=torch.channels_last)\n    if parallel:\n        model = enable_data_parallel(model)\n    return model\n\n\n\n\n\nclass EMA:\n    \"\"\"Exponential moving average over parameters *and* float buffers.\n\n    Including BatchNorm running statistics matters here: with augmentation this\n    strong, the raw running stats lag the EMA'd weights they are supposed to\n    normalise, which shows up as a validation AUC that trails the training\n    curve. Shadow keys are canonical, so the checkpoint stays portable.\"\"\"\n\n    def __init__(self, model, decay=None):\n        self.decay = config.EMA_DECAY if decay is None else decay\n        self.shadow = {}\n        self.backup = {}\n        for k, v in model.state_dict().items():\n            if torch.is_tensor(v) and v.is_floating_point():\n                self.shadow[canonical_key(k)] = v.detach().clone().float()\n\n    @torch.no_grad()\n    def update(self, model):\n        d = self.decay\n        for k, v in model.state_dict().items():\n            ck = canonical_key(k)\n            s = self.shadow.get(ck)\n            if s is not None and torch.is_tensor(v) and v.is_floating_point():\n                s.mul_(d).add_(v.detach().float(), alpha=1.0 - d)\n\n    @torch.no_grad()\n    def apply_shadow(self, model):\n        self.backup = {}\n        for k, v in model.state_dict().items():\n            ck = canonical_key(k)\n            s = self.shadow.get(ck)\n            if s is not None and torch.is_tensor(v) and v.is_floating_point():\n                self.backup[k] = v.detach().clone()\n                v.copy_(s.to(v.dtype))\n\n    @torch.no_grad()\n    def restore(self, model):\n        if not self.backup:\n            return\n        sd = model.state_dict()\n        for k, v in self.backup.items():\n            if k in sd:\n                sd[k].copy_(v)\n        self.backup = {}\n\n\nif __name__ == \"__main__\":\n    import sys as _sys\n    import types as _types\n    _cell_module = _types.ModuleType(\"model\")\n    _cell_module.__dict__.update(globals())\n    _sys.modules[\"model\"] = _cell_module\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport gc\nimport math\nimport time\nimport shutil\nimport argparse\nimport random\nimport threading\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom contextlib import contextmanager\nfrom sklearn.model_selection import StratifiedKFold, KFold\nfrom sklearn.metrics import roc_auc_score\n\nimport types\n\n\ndef _ensure_module(name, required=()):\n    \"\"\"Rebuild a project module from the shared notebook namespace.\n\n    When files are pasted as notebook cells, each cell shares one global\n    namespace. If an older pasted cell lacks the self-registration block, its\n    names still live in the namespace, so we can synthesize the module.\"\"\"\n    if name in sys.modules:\n        return sys.modules[name]\n    m = types.ModuleType(name)\n    m.__dict__.update(globals())\n    missing = [r for r in required if not hasattr(m, r)]\n    if missing:\n        raise ModuleNotFoundError(\n            f\"Module '{name}' is not importable and its cell has not been executed \"\n            f\"(missing: {', '.join(missing)}). Run the '{name}.py' cell before this one.\"\n        )\n    sys.modules[name] = m\n    return m\n\n\ntry:\n    import config\nexcept ModuleNotFoundError:\n    config = _ensure_module(\"config\", required=(\"TRAIN_CSV\", \"DEVICE\"))\n\ntry:\n    import data_utils as U\nexcept ModuleNotFoundError:\n    U = _ensure_module(\"data_utils\", required=(\"load_train_df\", \"get_report_paths\"))\n\ntry:\n    import dataset as D\nexcept ModuleNotFoundError:\n    D = _ensure_module(\"dataset\", required=(\"KneeDataset\", \"MultimodalCollator\",\n                                            \"build_vocab\"))\n\ntry:\n    from model import (build_model, EMA, build_criterion, compute_pos_weight,\n                       canonical_state_dict, load_canonical_state_dict)\nexcept ModuleNotFoundError:\n    _m = _ensure_module(\"model\", required=(\"build_model\", \"EMA\", \"build_criterion\"))\n    build_model, EMA = _m.build_model, _m.EMA\n    build_criterion = _m.build_criterion\n    compute_pos_weight = _m.compute_pos_weight\n    canonical_state_dict = _m.canonical_state_dict\n    load_canonical_state_dict = _m.load_canonical_state_dict\n\n\n\n\n\n_AMP_DTYPES = {\"float16\": torch.float16, \"fp16\": torch.float16,\n               \"bfloat16\": torch.bfloat16, \"bf16\": torch.bfloat16}\n\n\ndef amp_enabled():\n    return bool(config.USE_AMP and config.DEVICE == \"cuda\")\n\n\ndef amp_dtype():\n    \n    return _AMP_DTYPES.get(str(getattr(config, \"AMP_DTYPE\", \"float16\")).lower(),\n                           torch.float16)\n\n\ntry:\n    from torch.amp import autocast as _autocast, GradScaler as _GradScaler\n\n    def amp_ctx():\n        return _autocast(\"cuda\", dtype=amp_dtype(), enabled=amp_enabled())\n\n    def make_scaler():\n        return _GradScaler(\"cuda\", enabled=amp_enabled()\n                           and amp_dtype() == torch.float16)\nexcept ImportError:  \n    from torch.cuda.amp import autocast as _autocast, GradScaler as _GradScaler\n\n    def amp_ctx():\n        return _autocast(enabled=amp_enabled(), dtype=amp_dtype())\n\n    def make_scaler():\n        return _GradScaler(enabled=amp_enabled()\n                           and amp_dtype() == torch.float16)\n\n\ndef unwrap(model):\n    return model.module if isinstance(model, nn.DataParallel) else model\n\n\ndef _release_memory():\n    gc.collect()\n    if config.DEVICE == \"cuda\":\n        torch.cuda.empty_cache()\n\n\nclass _StallGuard:\n    \"\"\"Watchdog against silent DataParallel / DataLoader / shared-memory hangs.\n\n    If no iteration reports progress within HANG_TIMEOUT_S, print diagnostics\n    and hard-kill the process (which reaps every DataLoader worker), so a\n    deadlocked kernel cannot burn the session silently for hours.\"\"\"\n\n    def __init__(self, label, timeout_s=None):\n        self.label = label\n        self.timeout = (int(getattr(config, \"HANG_TIMEOUT_S\", 600))\n                        if timeout_s is None else int(timeout_s))\n        self.last = time.monotonic()\n        self._stop = threading.Event()\n        self._thread = None\n\n    def tick(self):\n        self.last = time.monotonic()\n\n    def start(self):\n        if self.timeout > 0:\n            self._thread = threading.Thread(target=self._watch, daemon=True)\n            self._thread.start()\n        return self\n\n    def stop(self):\n        self._stop.set()\n\n    def _watch(self):\n        poll = max(1, min(30, self.timeout // 4))\n        while not self._stop.wait(poll):\n            idle = time.monotonic() - self.last\n            if idle > self.timeout:\n                self._report(idle)\n                os._exit(137)\n\n    def _report(self, idle):\n        try:\n            lines = [\"=\" * 72,\n                     f\"[STALL] no progress for {idle / 60.0:.1f} min during \"\n                     f\"{self.label} — forcing cleanup\",\n                     \"=\" * 72]\n            if torch.cuda.is_available():\n                for g in range(torch.cuda.device_count()):\n                    try:\n                        lines.append(\n                            f\"[STALL] cuda:{g} allocated=\"\n                            f\"{torch.cuda.memory_allocated(g) / 2 ** 30:.2f} GB \"\n                            f\"reserved={torch.cuda.memory_reserved(g) / 2 ** 30:.2f} GB\")\n                    except Exception:\n                        pass\n            lines.append(\"[STALL] likely DataParallel / DataLoader / shared-memory \"\n                         \"deadlock. Re-run with DATA_PARALLEL=False and/or \"\n                         \"NUM_WORKERS=0 to isolate.\")\n            print(\"\\n\".join(lines), flush=True)\n        except Exception:\n            pass\n\n\n@contextmanager\ndef stall_guard(label, timeout_s=None):\n    g = _StallGuard(label, timeout_s).start()\n    try:\n        yield g\n    finally:\n        g.stop()\n\n\ndef set_seed(seed=None):\n    seed = config.SEED if seed is None else seed\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32 = True\n    if hasattr(torch, \"set_float32_matmul_precision\"):\n        torch.set_float32_matmul_precision(\"high\")\n\n\n\n\n\ndef _merge_rare_bins(strat, n_splits):\n    values, counts = np.unique(strat, return_counts=True)\n    if counts.min() >= n_splits:\n        return strat\n    ok_values = values[counts >= n_splits]\n    if len(ok_values) == 0:\n        return strat\n    remap = {}\n    for v, c in zip(values, counts):\n        remap[v] = v if c >= n_splits else \\\n            int(ok_values[np.argmin(np.abs(ok_values - v))])\n    merged = np.array([remap[s] for s in strat])\n    _, merged_counts = np.unique(merged, return_counts=True)\n    if merged_counts.min() < n_splits:\n        return strat\n    print(f\"[split] merged {int((merged != strat).sum())} rare strata rows into \"\n          f\"neighbouring bins for stable {n_splits}-fold splitting\")\n    return merged\n\n\ndef make_splits(df):\n    cols = U.get_target_columns(df)\n    if not cols:\n        kf = KFold(n_splits=config.FOLDS, shuffle=True, random_state=config.SEED)\n        return list(kf.split(df))\n    strat = df[cols].sum(axis=1).values.astype(int)\n    strat = _merge_rare_bins(strat, config.FOLDS)\n    try:\n        skf = StratifiedKFold(n_splits=config.FOLDS, shuffle=True,\n                              random_state=config.SEED)\n        return list(skf.split(df, strat))\n    except Exception:\n        kf = KFold(n_splits=config.FOLDS, shuffle=True, random_state=config.SEED)\n        return list(kf.split(df))\n\n\ndef per_class_auc(preds, targets, names=None):\n    names = names or config.TARGET_COLS\n    out = {}\n    for i in range(targets.shape[1]):\n        if len(np.unique(targets[:, i])) < 2:\n            continue\n        label = names[i] if i < len(names) else f\"t{i}\"\n        out[label] = float(roc_auc_score(targets[:, i], preds[:, i]))\n    return out\n\n\ndef macro_auc(preds, targets):\n    if preds is None or targets is None or len(preds) == 0:\n        return 0.0\n    aucs = list(per_class_auc(preds, targets).values())\n    return float(np.mean(aucs)) if aucs else 0.0\n\n\ndef smooth_labels(labels):\n    ls = config.LABEL_SMOOTH\n    return labels * (1.0 - ls) + 0.5 * ls\n\n\n\n\n\ndef mixup_batch(images, labels):\n    \"\"\"Kept for API compatibility: returns pre-blended labels.\"\"\"\n    lam = float(np.random.beta(config.MIXUP_ALPHA, config.MIXUP_ALPHA))\n    perm = torch.randperm(images.size(0), device=images.device)\n    return lam * images + (1.0 - lam) * images[perm], \\\n        lam * labels + (1.0 - lam) * labels[perm]\n\n\ndef _cutmix(images, lam):\n    \"\"\"Zero-copy-free spatial CutMix applied identically across all slices.\n\n    images is (B, N, C, H, W); the box is cut in (H, W) and shared by every\n    slice so the volume stays anatomically consistent.\"\"\"\n    b, _, _, h, w = images.shape\n    perm = torch.randperm(b, device=images.device)\n    ratio = math.sqrt(max(1.0 - lam, 0.0))\n    cut_h, cut_w = int(h * ratio), int(w * ratio)\n    if cut_h < 1 or cut_w < 1:\n        return images, perm, 1.0\n    cy, cx = random.randint(0, h - 1), random.randint(0, w - 1)\n    y1, y2 = max(cy - cut_h // 2, 0), min(cy + cut_h // 2, h)\n    x1, x2 = max(cx - cut_w // 2, 0), min(cx + cut_w // 2, w)\n    if y2 <= y1 or x2 <= x1:\n        return images, perm, 1.0\n    images = images.clone()\n    images[..., y1:y2, x1:x2] = images[perm][..., y1:y2, x1:x2]\n    real_lam = 1.0 - ((y2 - y1) * (x2 - x1) / float(h * w))\n    return images, perm, real_lam\n\n\ndef apply_mix(images, labels):\n    \"\"\"Return (images, target) where target is either labels or (y_a, y_b, lam).\n\n    Deferring the label blend to the loss keeps mixing compatible with losses\n    that are only defined on binary targets -- Asymmetric Loss and Focal both\n    assume y in {0,1}, so pre-blending the labels (as the previous code did)\n    silently changed what those losses optimise.\"\"\"\n    if images is None:\n        return images, labels\n    r = random.random()\n    if config.MIXUP_ALPHA > 0 and r < config.MIXUP_PROB:\n        lam = float(np.random.beta(config.MIXUP_ALPHA, config.MIXUP_ALPHA))\n        perm = torch.randperm(images.size(0), device=images.device)\n        images = lam * images + (1.0 - lam) * images[perm]\n        return images, (labels, labels[perm], lam)\n    if config.CUTMIX_ALPHA > 0 and r < config.MIXUP_PROB + config.CUTMIX_PROB:\n        lam = float(np.random.beta(config.CUTMIX_ALPHA, config.CUTMIX_ALPHA))\n        images, perm, real_lam = _cutmix(images, lam)\n        if real_lam >= 1.0:\n            return images, labels\n        return images, (labels, labels[perm], real_lam)\n    return images, labels\n\n\ndef mixed_loss(criterion, logits, target):\n    if isinstance(target, tuple):\n        y_a, y_b, lam = target\n        return lam * criterion(logits, y_a) + (1.0 - lam) * criterion(logits, y_b)\n    return criterion(logits, target)\n\n\n\n\n\ndef build_optimizer(model):\n    \"\"\"Discriminative LRs: a low rate for pretrained weights, a high one for the\n    randomly initialised depth-transformer, fusion and head. Norm/bias tensors\n    are excluded from weight decay.\"\"\"\n    decay, no_decay, bb_decay, bb_no_decay = [], [], [], []\n    for n, p in unwrap(model).named_parameters():\n        if not p.requires_grad:\n            continue\n        is_bb = n.startswith(\"backbone\")\n        skip_wd = p.ndim <= 1 or n.endswith(\".bias\") or \"pos\" in n.split(\".\")[-1]\n        if is_bb:\n            (bb_no_decay if skip_wd else bb_decay).append(p)\n        else:\n            (no_decay if skip_wd else decay).append(p)\n\n    groups = [\n        {\"params\": bb_decay, \"lr\": config.LR_BACKBONE, \"weight_decay\": config.WD},\n        {\"params\": bb_no_decay, \"lr\": config.LR_BACKBONE, \"weight_decay\": 0.0},\n        {\"params\": decay, \"lr\": config.LR, \"weight_decay\": config.WD},\n        {\"params\": no_decay, \"lr\": config.LR, \"weight_decay\": 0.0},\n    ]\n    return optim.AdamW([g for g in groups if g[\"params\"]], betas=(0.9, 0.999))\n\n\ndef build_scheduler(optimizer, total_steps, warmup_ratio=None, min_ratio=1e-3):\n    \"\"\"Linear warmup then cosine decay.\n\n    Written as a LambdaLR that clamps its own progress, so a skipped optimiser\n    step can never push it past the end -- OneCycleLR raises in that situation,\n    and any non-finite batch used to be enough to trigger it.\"\"\"\n    warmup_ratio = config.WARMUP_RATIO if warmup_ratio is None else warmup_ratio\n    total_steps = max(int(total_steps), 1)\n    warmup = max(1, int(round(total_steps * warmup_ratio)))\n\n    def fn(step):\n        if step < warmup:\n            return float(step + 1) / float(warmup)\n        p = (step - warmup) / float(max(total_steps - warmup, 1))\n        p = min(max(p, 0.0), 1.0)\n        return min_ratio + (1.0 - min_ratio) * 0.5 * (1.0 + math.cos(math.pi * p))\n\n    return optim.lr_scheduler.LambdaLR(optimizer, fn)\n\n\n\n\n\ndef _batch_to_device(batch, with_labels=True):\n    dev = config.DEVICE\n    img = batch.get(\"image\")\n    if img is not None:\n        img = img.to(dev, non_blocking=True)\n        \n        \n        img = img.to(amp_dtype() if amp_enabled() else torch.float32)\n    txt = batch.get(\"text_ids\")\n    if txt is not None:\n        txt = txt.to(dev, non_blocking=True)\n    msk = batch.get(\"text_mask\")\n    if msk is not None:\n        msk = msk.to(dev, non_blocking=True).float()\n    lab = batch.get(\"labels\") if with_labels else None\n    if lab is not None:\n        lab = lab.to(dev, non_blocking=True).float()\n    return img, txt, msk, lab\n\n\n\n\n\ndef train_one_epoch(model, ema, loader, optimizer, criterion, scaler, sched,\n                    accum, clip, log_every=0):\n    model.train()\n    optimizer.zero_grad(set_to_none=True)\n    total_loss, n_batch, n_skipped = 0.0, 0, 0\n    pending = 0          \n    n_total = len(loader)\n\n    guard = _StallGuard(\"training\").start()\n    for i, batch in enumerate(loader):\n        guard.tick()\n        img, txt, msk, lab = _batch_to_device(batch)\n        lab = smooth_labels(lab)\n        img, target = apply_mix(img, lab)\n\n        with amp_ctx():\n            logits = model(img, txt, msk)\n            loss = mixed_loss(criterion, logits, target)\n\n        step_due = ((i + 1) % accum == 0) or (i + 1 == n_total)\n\n        if torch.isfinite(loss):\n            scaler.scale(loss / accum).backward()\n            total_loss += float(loss.detach())\n            n_batch += 1\n            pending += 1\n        else:\n            n_skipped += 1\n\n        if not step_due:\n            continue\n\n        if pending == 0:\n            \n            \n            \n            optimizer.zero_grad(set_to_none=True)\n            sched.step()\n            continue\n\n        scaler.unscale_(optimizer)\n        \n        \n        grad_norm = nn.utils.clip_grad_norm_(\n            (p for p in model.parameters() if p.grad is not None), clip)\n        if not torch.isfinite(grad_norm):\n            n_skipped += 1\n        \n        \n        \n        \n        \n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad(set_to_none=True)\n        sched.step()\n        pending = 0\n        if ema is not None:\n            ema.update(unwrap(model))\n\n        if (i + 1) % max(1, int(getattr(config, \"GC_EVERY_BATCHES\", 64))) == 0:\n            _release_memory()\n\n        if log_every and ((i + 1) // accum) % log_every == 0:\n            _release_memory()\n            print(f\"    step {i + 1}/{n_total} loss={total_loss / max(n_batch, 1):.4f} \"\n                  f\"lr={optimizer.param_groups[-1]['lr']:.2e}\", flush=True)\n\n    guard.stop()\n    _release_memory()\n    if n_skipped:\n        print(f\"  [warn] skipped {n_skipped} non-finite batches/steps\",\n              flush=True)\n    return total_loss / max(n_batch, 1)\n\n\n_TTA_OPS = {\n    \"none\": lambda x: x,\n    \"hflip\": lambda x: x.flip(-1),\n    \"vflip\": lambda x: x.flip(-2),\n    \"shift\": lambda x: torch.roll(x, shifts=(6, 6), dims=(-2, -1)),\n    \"gamma\": lambda x: x.clamp(0.0, 1.0).pow(1.15),\n}\n\n\ndef _tta_ops(tta, has_image=True):\n    if not tta or not has_image:\n        \n        \n        return [\"none\"]\n    ops = [o for o in getattr(config, \"TTA_OPS\", [\"none\", \"hflip\"])\n           if o in _TTA_OPS]\n    return ops or [\"none\"]\n\n\n@torch.no_grad()\ndef evaluate(model, loader, tta=False, return_targets=True):\n    \"\"\"Predict over a loader under autocast, optionally with TTA.\n\n    torch.no_grad rather than inference_mode: DataParallel replicates the\n    backbone on every call and inference-mode tensors cannot be re-used across\n    replica boundaries on all torch versions.\"\"\"\n    model.eval()\n    ops = None\n    preds, targets = [], []\n    guard = _StallGuard(\"evaluation\").start()\n    for batch in loader:\n        guard.tick()\n        img, txt, msk, lab = _batch_to_device(batch)\n        if ops is None:\n            ops = _tta_ops(tta, has_image=img is not None)\n        with amp_ctx():\n            acc = None\n            for name in ops:\n                x = _TTA_OPS[name](img) if img is not None else None\n                p = torch.sigmoid(model(x, txt, msk).float())\n                acc = p if acc is None else acc + p\n            acc = acc / float(len(ops))\n        preds.append(acc.float().cpu().numpy())\n        if return_targets and lab is not None:\n            targets.append(lab.cpu().numpy())\n    guard.stop()\n    preds = np.vstack(preds) if preds else \\\n        np.zeros((0, config.NUM_TARGETS), np.float32)\n    return preds, (np.vstack(targets) if targets else None)\n\n\ndef evaluate_tta(model, loader, tta):\n    \"\"\"Backwards-compatible wrapper used by earlier revisions.\"\"\"\n    return evaluate(model, loader, tta=tta, return_targets=False)\n\n\n\n\n\ndef _fold_vocab(train_df, reports, fallback):\n    \"\"\"Vocabulary from this fold's training rows only.\n\n    Building it from the full dataframe leaked validation-fold token statistics\n    into every model. It is stored in the checkpoint so inference reproduces the\n    exact mapping each fold was trained with.\"\"\"\n    if not getattr(config, \"FOLD_VOCAB\", True) or not reports:\n        return fallback\n    return D.build_vocab(train_df, reports)\n\n\ndef run_fold(fold, train_idx, valid_idx, df, reports, vocab, epochs,\n             cache_dir=None):\n    set_seed(config.SEED + fold)\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    valid_df = df.iloc[valid_idx].reset_index(drop=True)\n    fold_vocab = _fold_vocab(train_df, reports, vocab)\n\n    tr_ds = D.KneeDataset(train_df, reports, fold_vocab, mode=\"train\",\n                          volume_root=config.VOLUME_ROOT, cache_dir=cache_dir)\n    va_ds = D.KneeDataset(valid_df, reports, fold_vocab, mode=\"valid\",\n                          volume_root=config.VOLUME_ROOT, cache_dir=cache_dir)\n    tr_loader = D.make_loader(tr_ds, shuffle=True, drop_last=True)\n    va_loader = D.make_loader(va_ds, shuffle=False)\n\n    model = build_model(vocab_size=len(fold_vocab))\n    ema = EMA(model) if config.EMA_DECAY else None\n\n    optimizer = build_optimizer(model)\n    steps_per_epoch = max(math.ceil(len(tr_loader) / config.ACCUM_STEPS), 1)\n    sched = build_scheduler(optimizer, epochs * steps_per_epoch)\n\n    cols_tr = U.get_target_columns(train_df)\n    criterion = build_criterion(\n        train_df[cols_tr].values.astype(np.float32) if cols_tr else None)\n    scaler = make_scaler()\n\n    best_auc, best_epoch, best_preds = 0.0, -1, None\n    for ep in range(1, epochs + 1):\n        print(f\"[Fold {fold}] epoch {ep}/{epochs} starting\", flush=True)\n        loss = train_one_epoch(model, ema, tr_loader, optimizer, criterion,\n                               scaler, sched, config.ACCUM_STEPS,\n                               config.GRAD_CLIP,\n                               log_every=getattr(config, \"LOG_EVERY_STEPS\", 10))\n        \n        \n        \n        if ema is not None:\n            ema.apply_shadow(unwrap(model))\n        preds, targets = evaluate(model, va_loader, tta=False)\n        if ema is not None:\n            ema.restore(unwrap(model))\n        auc = macro_auc(preds, targets)\n        print(f\"[Fold {fold}] epoch {ep}/{epochs} loss={loss:.4f} \"\n              f\"val_mAUC={auc:.5f}\", flush=True)\n\n        if auc > best_auc:\n            best_auc, best_epoch, best_preds = auc, ep, preds\n            torch.save({\"model\": canonical_state_dict(unwrap(model)),\n                        \"ema\": ema.shadow if ema else None,\n                        \"auc\": auc, \"epoch\": ep,\n                        \"backbone\": config.MODEL_NAME,\n                        \"img_size\": config.IMG_SIZE,\n                        \"n_slices\": config.N_SEQ_IMAGES,\n                        \"vocab\": fold_vocab,\n                        \"vocab_size\": len(fold_vocab)},\n                       os.path.join(config.MODEL_DIR, f\"fold{fold}.pt\"))\n        _release_memory()\n\n    if best_preds is not None and targets is not None:\n        worst = sorted(per_class_auc(best_preds, targets).items(),\n                       key=lambda kv: kv[1])[:3]\n        print(f\"[Fold {fold}] weakest classes: \"\n              + \", \".join(f\"{k}={v:.3f}\" for k, v in worst), flush=True)\n    print(f\"[Fold {fold}] best val mAUC={best_auc:.5f} @ ep{best_epoch}\",\n          flush=True)\n\n    del model, ema, optimizer, tr_loader, va_loader\n    _release_memory()\n    return best_auc, best_epoch, best_preds\n\n\n\n\n\ndef _strip(state):\n    return {k.replace(\"module.\", \"\"): v for k, v in state.items()}\n\n\ndef _parse_folds(folds):\n    if isinstance(folds, str) and \",\" in folds:\n        return [int(x) for x in folds.split(\",\") if x.strip()] or [0]\n    if isinstance(folds, (list, tuple, range)):\n        return [int(x) for x in folds]\n    if isinstance(folds, str) and folds.isdigit():\n        return [int(folds)]\n    if isinstance(folds, int):\n        return [folds]\n    return [0]\n\n\ndef _test_volume_root():\n    te = getattr(config, \"VOLUME_TEST_ROOT\", None)\n    if te and os.path.isdir(te):\n        return te\n    return getattr(config, \"VOLUME_ROOT\", None)\n\n\ndef rank_normalize(probs):\n    \"\"\"Column-wise rank transform to [0, 1].\n\n    AUC only reads the ordering within a column, so ranks are the natural space\n    to average folds in: a fold that is well ordered but poorly calibrated no\n    longer drags the blend toward its own probability scale.\"\"\"\n    return pd.DataFrame(probs).rank(axis=0, pct=True,\n                                    method=\"average\").values.astype(np.float64)\n\n\ndef ensemble_folds(fold_probs, weights=None, mode=None):\n    \"\"\"Blend per-fold predictions. mode: 'rank' | 'auc_weighted' | 'mean'.\"\"\"\n    mode = (mode or getattr(config, \"ENSEMBLE\", \"rank\")).lower()\n    if not fold_probs:\n        raise ValueError(\"nothing to ensemble\")\n    n_rows = fold_probs[0].shape[0]\n    w = np.ones(len(fold_probs)) if weights is None else \\\n        np.asarray(weights, dtype=np.float64)\n    if w.sum() <= 0:\n        w = np.ones(len(fold_probs))\n    w = w / w.sum()\n\n    if mode == \"rank\":\n        if n_rows < 8:\n            \n            \n            print(f\"[predict] only {n_rows} rows; using weighted mean \"\n                  f\"instead of rank averaging\")\n            mode = \"auc_weighted\"\n        else:\n            return sum(wi * rank_normalize(p) for wi, p in zip(w, fold_probs))\n    if mode == \"mean\":\n        w = np.ones(len(fold_probs)) / len(fold_probs)\n    return sum(wi * p.astype(np.float64) for wi, p in zip(w, fold_probs))\n\n\ndef predict(df, reports, vocab, folds=\"all\", tta=None, volume_root=None,\n            cache_dir=None, mode=None):\n    \"\"\"Predict for `df` by ensembling every available fold checkpoint.\"\"\"\n    if tta is None:\n        tta = config.TTA\n    fold_list = list(range(config.FOLDS)) if folds == \"all\" else _parse_folds(folds)\n    if volume_root is None:\n        volume_root = _test_volume_root()\n\n    fold_probs, fold_w, used = [], [], []\n    loader_cache = {}\n    model = None\n\n    for fold in fold_list:\n        ckpt_path = os.path.join(config.MODEL_DIR, f\"fold{fold}.pt\")\n        if not os.path.exists(ckpt_path):\n            print(f\"[predict] missing {ckpt_path}; skipping fold {fold}\")\n            continue\n        ckpt = torch.load(ckpt_path, map_location=\"cpu\")\n        fold_vocab = ckpt.get(\"vocab\") or vocab\n\n        \n        \n        \n        \n        key = (len(fold_vocab), tuple(list(fold_vocab)[:32]))\n        if key not in loader_cache:\n            ds = D.KneeDataset(df, reports, fold_vocab, mode=\"valid\",\n                               volume_root=volume_root, cache_dir=cache_dir)\n            loader_cache[key] = D.make_loader(ds, shuffle=False)\n        loader = loader_cache[key]\n\n        if model is None or model_vocab_size != len(fold_vocab):\n            del model\n            _release_memory()\n            model = build_model(vocab_size=len(fold_vocab))\n            model_vocab_size = len(fold_vocab)\n\n        state = dict(ckpt[\"model\"])\n        for k, v in (ckpt.get(\"ema\") or {}).items():\n            if k in state:\n                state[k] = v\n        load_canonical_state_dict(unwrap(model), state, verbose=False)\n\n        preds, _ = evaluate(model, loader, tta=tta, return_targets=False)\n        fold_probs.append(preds.astype(np.float64))\n        fold_w.append(float(ckpt.get(\"auc\", 0.0)) or 1e-3)\n        used.append(fold)\n        del ckpt, state, preds\n        _release_memory()\n\n    if not used:\n        raise FileNotFoundError(f\"no fold checkpoints found in {config.MODEL_DIR}; \"\n                                f\"train first (ACTIONS = ['--folds', 'all', ...])\")\n    mode = mode or getattr(config, \"ENSEMBLE\", \"rank\")\n    print(f\"[predict] ensembled folds {used} via '{mode}' \"\n          f\"(TTA={_tta_ops(tta)})\")\n    del model, loader_cache\n    _release_memory()\n    return ensemble_folds(fold_probs, fold_w, mode=mode)\n\n\ndef _publish_best_model(best_fold):\n    \"\"\"Copy the best fold's checkpoint to /kaggle/working/best_auc_model.pth.\"\"\"\n    src = os.path.join(config.MODEL_DIR, f\"fold{best_fold}.pt\")\n    if not os.path.isfile(src):\n        print(f\"[model] no checkpoint to publish for fold {best_fold}\", flush=True)\n        return None\n    try:\n        os.makedirs(config.WORKING_DIR, exist_ok=True)\n        dst = os.path.join(config.WORKING_DIR, \"best_auc_model.pth\")\n        shutil.copy2(src, dst)\n        print(f\"[model] best fold {best_fold} checkpoint -> {dst}\", flush=True)\n        return dst\n    except Exception as e:\n        print(f\"[model] could not publish best checkpoint: {e}\", flush=True)\n        return None\n\n\ndef write_submission(tdf, probs):\n    sub = pd.DataFrame({\"StudyInstanceUID\": tdf[U.uid_column(tdf)].values})\n    for i, c in enumerate(config.TARGET_COLS):\n        sub[c] = probs[:, i]\n    os.makedirs(config.WORKING_DIR, exist_ok=True)\n    out = os.path.join(config.WORKING_DIR, \"submission.csv\")\n    sub.to_csv(out, index=False)\n    print(\"Saved\", out, flush=True)\n    return out\n\n\n\n\n\ndef prepare_cache(df, test_df=None):\n    \"\"\"Decode every study to the uint8 mmap cache once, up front.\"\"\"\n    if not getattr(config, \"USE_VOLUME_CACHE\", True):\n        return None\n    if getattr(config, \"CLEAN_CACHE_ON_START\", False):\n        U.clear_volume_cache(config.VOLUME_CACHE_DIR)\n    jobs = [(df, config.VOLUME_ROOT)]\n    if test_df is not None:\n        jobs.append((test_df, _test_volume_root()))\n    info = None\n    for frame, root in jobs:\n        if frame is None or not root:\n            continue\n        info = U.build_volume_cache([str(u) for u in frame[U.uid_column(frame)]],\n                                    volume_root=root,\n                                    cache_dir=config.VOLUME_CACHE_DIR)\n    if not info or not info.get(\"cached\"):\n        return None\n    return config.VOLUME_CACHE_DIR\n\n\ndef _audit_text_modality(coverage, reports, df):\n    \"\"\"Report how much text will actually exist at inference time.\n\n    The reports restate the findings almost verbatim, so a fusion head trained\n    on full-coverage text and served none of it collapses. When coverage is\n    poor, modality dropout is raised so the vision branch has to carry the\n    prediction on its own.\"\"\"\n    uid_col = U.uid_column(df)\n    train_hit = sum(1 for u in df[uid_col] if reports.get(str(u), \"\").strip())\n    train_cov = train_hit / float(max(len(df), 1))\n    print(f\"[data] report coverage: train={train_cov:.1%}  test={coverage:.1%}\")\n    if train_cov > 0.5 and coverage < 0.5:\n        was = config.MODALITY_DROPOUT_TEXT\n        config.MODALITY_DROPOUT_TEXT = max(was, 0.5)\n        print(f\"[data] WARNING: reports exist for training but are largely \"\n              f\"absent at test time. Raising MODALITY_DROPOUT_TEXT \"\n              f\"{was:.2f} -> {config.MODALITY_DROPOUT_TEXT:.2f} so the vision \"\n              f\"branch stays predictive on its own.\")\n    if train_cov < 0.02:\n        print(\"[data] text modality is effectively empty; the model is \"\n              \"vision-only this run.\")\n\n\ndef main():\n    for _stream in (sys.stdout, sys.stderr):\n        try:\n            _stream.reconfigure(line_buffering=True)\n        except Exception:\n            pass\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--folds\", default=\"all\")\n    ap.add_argument(\"--epochs\", type=int, default=config.EPOCHS)\n    ap.add_argument(\"--predict\", action=\"store_true\")\n    ap.add_argument(\"--submit\", action=\"store_true\")\n    ap.add_argument(\"--prepare\", action=\"store_true\",\n                    help=\"build the volume cache and exit\")\n    ap.add_argument(\"--no-cache\", action=\"store_true\",\n                    help=\"read DICOM directly instead of the mmap cache\")\n    args = ap.parse_args()\n\n    set_seed()\n    D.configure_sharing_strategy()\n    if args.no_cache:\n        config.USE_VOLUME_CACHE = False\n    if getattr(config, \"CLEAN_WORKING_ON_START\", True):\n        U.clear_working_cache()\n\n    df = U.load_train_df()\n    try:\n        tdf = U.load_test_df()\n    except Exception as e:\n        print(f\"[data] test set unavailable ({e})\")\n        tdf = None\n\n    reports_text, coverage = U.collect_reports(df, tdf)\n    if not reports_text:\n        print(\"[data] WARNING: no report files and no report column; \"\n              \"text modality disabled\")\n    else:\n        _audit_text_modality(coverage, reports_text, df)\n\n    cache_dir = prepare_cache(df, tdf)\n    if args.prepare:\n        return\n\n    vocab = D.build_vocab(df, reports_text)\n    cols = U.get_target_columns(df)\n\n    if not args.predict:\n        splits = make_splits(df)\n        fold_list = (list(range(len(splits))) if args.folds == \"all\"\n                     else _parse_folds(args.folds))\n        oof = np.full((len(df), config.NUM_TARGETS), np.nan, dtype=np.float32)\n        aucs = []\n        for f in fold_list:\n            tr, va = splits[f]\n            a, ep, preds = run_fold(f, tr, va, df, reports_text, vocab,\n                                    args.epochs, cache_dir=cache_dir)\n            if preds is not None:\n                oof[va] = preds\n            aucs.append(a)\n        print(f\"CV avg mAUC = {np.mean(aucs):.5f}\", flush=True)\n        best_fold = fold_list[int(np.argmax(aucs))]\n        _publish_best_model(best_fold)\n        os.makedirs(config.OUTPUT_DIR, exist_ok=True)\n        np.save(os.path.join(config.OUTPUT_DIR, \"oof.npy\"), oof)\n        if cols:\n            idx = [config.TARGET_COLS.index(c) for c in cols]\n            mask = ~np.isnan(oof[:, idx[0]])\n            oof_auc = macro_auc(oof[mask][:, idx], df.loc[mask, cols].values)\n            print(f\"Honest OOF mAUC = {oof_auc:.5f}\")\n            for k, v in sorted(per_class_auc(oof[mask][:, idx],\n                                             df.loc[mask, cols].values,\n                                             names=cols).items(),\n                               key=lambda kv: kv[1]):\n                print(f\"  {k:<18} {v:.4f}\")\n\n    if args.predict or args.submit:\n        if tdf is None:\n            print(\"[warn] test set unavailable; skipping submission\")\n            return\n        te = predict(tdf, reports_text, vocab, folds=args.folds, tta=config.TTA,\n                     cache_dir=cache_dir)\n        write_submission(tdf, te)\n\n\n\nbatched_train_one_epoch = train_one_epoch\n\nif __name__ == \"__main__\" and \"ipykernel\" not in sys.modules:\n    main()\n\nif __name__ == \"__main__\":\n    import types as _types\n    _cell_module = _types.ModuleType(\"train\")\n    _cell_module.__dict__.update(globals())\n    sys.modules[\"train\"] = _cell_module\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport os\nimport config\n\ntry:\n    from train import main\nexcept ModuleNotFoundError:\n    main = globals().get(\"main\")\n    if main is None:\n        raise\n\n\nACTIONS = [\"--folds\", \"all\", \"--epochs\", \"10\", \"--submit\"]\n\nos.makedirs(config.MODEL_DIR, exist_ok=True)\nsys.argv = [\"train.py\"] + ACTIONS\nmain()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}