{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"499388e3","cell_type":"markdown","source":"# Explainable Hybrid CNN-Transformer for Diabetic Retinopathy Grading\n\nThis notebook implements a **research-grounded 5-class diabetic retinopathy grading pipeline** with:\n\n- **APTOS 2019** as the primary Kaggle competition dataset\n- a **hybrid EfficientNet-B3 + DeiT-Tiny** fusion model\n- **ordinal-aware training** to reduce adjacent-grade confusion\n- **hard-example emphasis** during optimization\n- a held-out **calibration split** with **temperature scaling** before uncertainty-based referral\n- **confidence-based referral** for uncertain borderline cases using calibrated probabilities\n- **Grad-CAM** and **Attention Rollout** for explainability\n- **quantitative XAI validation** with faithfulness and stability checks\n\nThe implementation is intentionally written from scratch for this DR task and does not reuse the structure or code of the inspiration notebooks. The contribution should be described as a **task-specific integration and evaluation of established methods**, not as a claim of inventing a brand-new foundation architecture.","metadata":{}},{"id":"c3c54414","cell_type":"markdown","source":"## Notebook Roadmap\n\n1. Configure dataset paths and runtime\n2. Build a training dataframe from APTOS\n3. Create train, calibration, and validation splits\n4. Train a custom hybrid DR grader with live training widgets\n5. Fit temperature scaling on a held-out calibration split\n6. Evaluate classification quality with QWK, macro F1, AUC, ECE, and confusion analysis\n7. Add referral logic for uncertain adjacent-grade predictions using calibrated probabilities\n8. Generate and validate explanations with Grad-CAM and Attention Rollout\n9. Export APTOS test predictions as a submission file","metadata":{}},{"id":"765d8bc3","cell_type":"markdown","source":"## Research-Grounded Contribution Framing\n\nThis notebook is best framed as a **system-style DR grading study** rather than a claim of architectural novelty. The main contribution is the careful integration of:\n\n- **EfficientNet**-style CNN features for local lesion cues\n- **DeiT**-style transformer features for wider retinal context\n- **ordinal-aware supervision** for the ordered 5-grade DR label space\n- **temperature-scaled calibration** before threshold-based referral\n- **dual explainability** with Grad-CAM and attention rollout\n\nThese choices are aligned with well-known literature, including EfficientNet (Tan and Le, 2019), DeiT (Touvron et al., 2021), Grad-CAM (Selvaraju et al., 2017), and temperature scaling for neural calibration (Guo et al., 2017).","metadata":{}},{"id":"e511ee7f","cell_type":"code","source":"import importlib\nimport subprocess\nimport sys\n\n\ndef ensure_package(module_name: str, pip_name: str | None = None) -> None:\n    if importlib.util.find_spec(module_name) is None:\n        subprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", pip_name or module_name])\n\n\nfor module_name, pip_name in [\n    (\"timm\", \"timm\"),\n    (\"ipywidgets\", \"ipywidgets\"),\n]:\n    ensure_package(module_name, pip_name)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:51.150384Z","iopub.execute_input":"2026-04-09T08:50:51.151041Z","iopub.status.idle":"2026-04-09T08:50:51.156240Z","shell.execute_reply.started":"2026-04-09T08:50:51.151009Z","shell.execute_reply":"2026-04-09T08:50:51.155443Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d905951d","cell_type":"code","source":"from __future__ import annotations\n\nimport copy\nimport gc\nimport json\nimport os\nimport random\nimport shutil\nimport zipfile\nfrom contextlib import nullcontext\nfrom dataclasses import asdict, dataclass\nfrom pathlib import Path\n\nKAGGLE_CPU_LIMIT = 4\nAVAILABLE_CPU_COUNT = os.cpu_count() or 1\nCPU_WORKERS = min(KAGGLE_CPU_LIMIT, AVAILABLE_CPU_COUNT)\nfor env_var in (\n    \"OMP_NUM_THREADS\",\n    \"OPENBLAS_NUM_THREADS\",\n    \"MKL_NUM_THREADS\",\n    \"NUMEXPR_NUM_THREADS\",\n    \"VECLIB_MAXIMUM_THREADS\",\n):\n    os.environ[env_var] = str(CPU_WORKERS)\n\nimport cv2\nimport ipywidgets as widgets\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom IPython.display import FileLink, display\nfrom PIL import Image\nfrom sklearn.metrics import (\n    accuracy_score,\n    balanced_accuracy_score,\n    classification_report,\n    cohen_kappa_score,\n    confusion_matrix,\n    f1_score,\n    precision_recall_fscore_support,\n    roc_curve,\n    roc_auc_score,\n)\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom sklearn.preprocessing import label_binarize\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom torchvision import transforms\n\ntorch.set_num_threads(CPU_WORKERS)\nif hasattr(torch, \"set_num_interop_threads\"):\n    try:\n        torch.set_num_interop_threads(CPU_WORKERS)\n    except RuntimeError:\n        pass\ncv2.setNumThreads(CPU_WORKERS)\ncv2.setUseOptimized(True)\n\nsns.set_theme(style=\"whitegrid\")\n\nCLASS_NAMES = [\"No_DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative_DR\"]\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\n\n@dataclass\nclass Config:\n    aptos_root: str = \"/kaggle/input/competitions/aptos2019-blindness-detection\"\n    working_root: str = \"/kaggle/working/dr_hybrid_project\"\n    image_size: int = 224\n    batch_size: int = 32\n    num_classes: int = 5\n    num_workers: int = CPU_WORKERS\n    prefetch_factor: int = 4\n    epochs: int = 30\n    backbone_lr: float = 1.0e-4\n    head_lr: float = 3.0e-4\n    min_lr: float = 5e-6\n    warmup_epochs: int = 4\n    early_stopping_patience: int = 10\n    min_epochs_before_early_stop: int = 14\n    max_grad_norm: float = 1.0\n    weight_decay: float = 2.0e-4\n    head_weight_decay: float = 3.0e-4\n    label_smoothing: float = 0.04\n    val_size: float = 0.2\n    calibration_size: float = 0.1\n    seed: int = 42\n    deterministic: bool = False\n    pretrained_backbones: bool = True\n    cnn_backbone: str = \"efficientnet_b3\"\n    vit_backbone: str = \"deit_tiny_patch16_224\"\n    freeze_backbone_epochs: int = 4\n    class_balance_power: float = 0.70\n    sampler_weight_cap: float = 2.40\n    loss_class_balance_power: float = 0.35\n    loss_weight_floor: float = 0.85\n    loss_weight_cap: float = 1.75\n    use_ema: bool = True\n    ema_decay: float = 0.9992\n    tta_horizontal: bool = True\n    apply_clahe: bool = True\n    cache_preprocessed_images: bool = True\n    use_weighted_sampler: bool = True\n    use_loss_class_weights: bool = True\n    use_channels_last: bool = True\n    use_training_widgets: bool = True\n    ordinal_weight: float = 0.45\n    distance_weight: float = 0.18\n    hard_example_gamma: float = 1.05\n    advanced_class_start: int = 3\n    advanced_class_boost: float = 0.40\n    undergrading_penalty_weight: float = 0.35\n    undergrading_margin_power: float = 1.50\n    gradcam_samples: int = 5\n    xai_eval_samples: int = 15\n    referral_confidence_threshold: float = 0.60\n    referral_margin_threshold: float = 0.12\n    referral_min_coverage: float = 0.65\n    deletion_fraction: float = 0.15\n    stability_noise_std: float = 0.01\n    temperature_min: float = 0.5\n    temperature_max: float = 3.0\n    enable_threshold_optimization: bool = True\n    threshold_search_steps: int = 61\n    threshold_search_passes: int = 4\n    threshold_focus_upper_boundaries_only: bool = True\n    threshold_fixed_lower_boundaries: int = 2\n    threshold_search_radius: float = 0.45\n    threshold_max_shift: float = 0.55\n    threshold_min_gap: float = 0.18\n    threshold_max_accuracy_drop: float = 0.015\n    threshold_max_qwk_drop: float = 0.010\n    threshold_max_ece_increase: float = 0.020\n    threshold_min_score_gain: float = 0.002\n    threshold_min_advanced_recall_gain: float = 0.040\n    threshold_min_min_recall_gain: float = 0.030\n\n\nCFG = Config()\nCFG.num_workers = min(int(CFG.num_workers), CPU_WORKERS)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = torch.cuda.is_available()\nUSE_CHANNELS_LAST = DEVICE.type == \"cuda\" and CFG.use_channels_last\nif DEVICE.type == \"cuda\":\n    if hasattr(torch.backends.cuda.matmul, \"allow_tf32\"):\n        torch.backends.cuda.matmul.allow_tf32 = True\n    if hasattr(torch.backends.cudnn, \"allow_tf32\"):\n        torch.backends.cudnn.allow_tf32 = True\nif hasattr(torch, \"set_float32_matmul_precision\"):\n    torch.set_float32_matmul_precision(\"high\")\nWORKDIR = Path(CFG.working_root)\nWORKDIR.mkdir(parents=True, exist_ok=True)\nARTIFACT_DIR = WORKDIR / \"notebook_artifacts\"\nFIGURE_DIR = ARTIFACT_DIR / \"figures\"\nTABLE_DIR = ARTIFACT_DIR / \"tables\"\nREPORT_DIR = ARTIFACT_DIR / \"reports\"\nfor output_dir in (ARTIFACT_DIR, FIGURE_DIR, TABLE_DIR, REPORT_DIR):\n    output_dir.mkdir(parents=True, exist_ok=True)\nEXPORTED_ARTIFACTS: list[Path] = []\n\nprint(\"Device:\", DEVICE)\nprint(\"Available CPU cores:\", AVAILABLE_CPU_COUNT)\nprint(\"Configured CPU workers:\", CFG.num_workers)\nprint(\"Torch CPU threads:\", torch.get_num_threads())\nprint(\"OpenCV threads:\", cv2.getNumThreads())\nprint(\"Mixed precision:\", USE_AMP)\nprint(\"Channels-last tensors:\", USE_CHANNELS_LAST)\nprint(\"Working directory:\", WORKDIR)\nprint(asdict(CFG))","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:51.157757Z","iopub.execute_input":"2026-04-09T08:50:51.158077Z","iopub.status.idle":"2026-04-09T08:50:51.203016Z","shell.execute_reply.started":"2026-04-09T08:50:51.158035Z","shell.execute_reply":"2026-04-09T08:50:51.202180Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"9c4d3779","cell_type":"code","source":"def seed_everything(seed: int, deterministic: bool = False) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = deterministic\n    torch.backends.cudnn.benchmark = not deterministic\n\n\nseed_everything(CFG.seed, deterministic=CFG.deterministic)\n\n\ndef register_artifact(path: str | Path) -> Path:\n    artifact_path = Path(path)\n    if artifact_path not in EXPORTED_ARTIFACTS:\n        EXPORTED_ARTIFACTS.append(artifact_path)\n    return artifact_path\n\n\ndef export_json_artifact(payload: dict, filename: str) -> Path:\n    artifact_path = REPORT_DIR / filename\n    artifact_path.write_text(json.dumps(payload, indent=2), encoding=\"utf-8\")\n    return register_artifact(artifact_path)\n\n\ndef export_text_artifact(text: str, filename: str, directory: Path | None = None) -> Path:\n    target_dir = REPORT_DIR if directory is None else directory\n    artifact_path = target_dir / filename\n    artifact_path.write_text(text, encoding=\"utf-8\")\n    return register_artifact(artifact_path)\n\n\ndef export_table_bundle(\n    frame: pd.DataFrame,\n    stem: str,\n    include_index: bool = False,\n    round_digits: int = 4,\n) -> list[Path]:\n    artifact_paths = []\n\n    csv_path = TABLE_DIR / f\"{stem}.csv\"\n    frame.to_csv(csv_path, index=include_index)\n    artifact_paths.append(register_artifact(csv_path))\n\n    rounded_frame = frame.copy()\n    numeric_columns = rounded_frame.select_dtypes(include=[np.number]).columns\n    if len(numeric_columns):\n        rounded_frame[numeric_columns] = rounded_frame[numeric_columns].round(round_digits)\n\n    latex_path = TABLE_DIR / f\"{stem}.tex\"\n    latex_path.write_text(\n        rounded_frame.to_latex(index=include_index, escape=False),\n        encoding=\"utf-8\",\n    )\n    artifact_paths.append(register_artifact(latex_path))\n    return artifact_paths\n\n\ndef save_figure(fig, filename: str, dpi: int = 300) -> Path:\n    figure_path = FIGURE_DIR / filename\n    fig.savefig(figure_path, dpi=dpi, bbox_inches=\"tight\")\n    return register_artifact(figure_path)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:51.221936Z","iopub.execute_input":"2026-04-09T08:50:51.222132Z","iopub.status.idle":"2026-04-09T08:50:51.233882Z","shell.execute_reply.started":"2026-04-09T08:50:51.222113Z","shell.execute_reply":"2026-04-09T08:50:51.232954Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e4e4482b","cell_type":"markdown","source":"## Data Preparation\n\nAPTOS is loaded directly from the competition CSV and image folders. The training pipeline in this notebook uses the APTOS dataset only.","metadata":{}},{"id":"7a5651cf","cell_type":"code","source":"def verify_aptos_layout(cfg: Config) -> None:\n    expected_paths = [\n        Path(cfg.aptos_root) / \"train.csv\",\n        Path(cfg.aptos_root) / \"test.csv\",\n        Path(cfg.aptos_root) / \"train_images\",\n        Path(cfg.aptos_root) / \"test_images\",\n    ]\n    for path in expected_paths:\n        if not path.exists():\n            raise FileNotFoundError(f\"Missing expected APTOS path: {path}\")\n\n\ndef read_csv_from_zip(zip_path: Path) -> pd.DataFrame:\n    with zipfile.ZipFile(zip_path) as archive:\n        csv_candidates = [name for name in archive.namelist() if name.lower().endswith(\".csv\")]\n        if not csv_candidates:\n            raise FileNotFoundError(f\"No CSV file found inside {zip_path}\")\n        with archive.open(csv_candidates[0]) as handle:\n            return pd.read_csv(handle)\n\n\ndef join_multipart_zip(root: Path, zip_stem: str, output_dir: Path) -> Path:\n    output_dir.mkdir(parents=True, exist_ok=True)\n    assembled_zip = output_dir / f\"{zip_stem}.zip\"\n    if assembled_zip.exists():\n        return assembled_zip\n\n    split_parts = sorted(root.glob(f\"{zip_stem}.zip.*\"))\n    standalone_zip = root / f\"{zip_stem}.zip\"\n    if standalone_zip.exists():\n        return standalone_zip\n    if not split_parts:\n        raise FileNotFoundError(f\"No archive parts found for {zip_stem} under {root}\")\n\n    with assembled_zip.open(\"wb\") as target_handle:\n        for part in split_parts:\n            with part.open(\"rb\") as source_handle:\n                shutil.copyfileobj(source_handle, target_handle)\n    return assembled_zip\n\n\ndef extract_zip_if_needed(zip_path: Path, extract_to: Path) -> Path:\n    marker = extract_to / \".complete\"\n    if marker.exists():\n        return extract_to\n\n    extract_to.mkdir(parents=True, exist_ok=True)\n    with zipfile.ZipFile(zip_path) as archive:\n        archive.extractall(extract_to)\n    marker.touch()\n    return extract_to\n\n\ndef crop_black_border(image_rgb: np.ndarray, tolerance: int = 7) -> np.ndarray:\n    grayscale = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2GRAY)\n    mask = grayscale > tolerance\n    if mask.sum() == 0:\n        return image_rgb\n    coordinates = np.argwhere(mask)\n    y0, x0 = coordinates.min(axis=0)\n    y1, x1 = coordinates.max(axis=0) + 1\n    return image_rgb[y0:y1, x0:x1]\n\n\ndef load_retina_image(image_path: str | Path, image_size: int, apply_clahe: bool = False) -> Image.Image:\n    image_path = Path(image_path)\n    image_bgr = cv2.imread(str(image_path))\n    if image_bgr is None:\n        raise FileNotFoundError(f\"Could not read image: {image_path}\")\n    image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)\n    image_rgb = crop_black_border(image_rgb)\n\n    if apply_clahe:\n        lab = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2LAB)\n        l_channel, a_channel, b_channel = cv2.split(lab)\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        l_channel = clahe.apply(l_channel)\n        image_rgb = cv2.cvtColor(cv2.merge((l_channel, a_channel, b_channel)), cv2.COLOR_LAB2RGB)\n\n    image_rgb = cv2.resize(image_rgb, (image_size, image_size), interpolation=cv2.INTER_AREA)\n    return Image.fromarray(image_rgb)\n\n\ndef build_aptos_train_frame(cfg: Config) -> pd.DataFrame:\n    verify_aptos_layout(cfg)\n    frame = pd.read_csv(Path(cfg.aptos_root) / \"train.csv\")\n    frame = frame.rename(columns={\"id_code\": \"image_id\", \"diagnosis\": \"label\"})\n    frame[\"image_path\"] = frame[\"image_id\"].map(lambda name: str(Path(cfg.aptos_root) / \"train_images\" / f\"{name}.png\"))\n    frame[\"source\"] = \"APTOS2019\"\n    frame[\"label\"] = frame[\"label\"].astype(int)\n    frame = frame[frame[\"image_path\"].map(lambda path: Path(path).exists())].reset_index(drop=True)\n    return frame[[\"image_id\", \"image_path\", \"label\", \"source\"]]\n\n\ndef build_aptos_test_frame(cfg: Config) -> pd.DataFrame:\n    frame = pd.read_csv(Path(cfg.aptos_root) / \"test.csv\")\n    frame = frame.rename(columns={\"id_code\": \"image_id\"})\n    frame[\"image_path\"] = frame[\"image_id\"].map(lambda name: str(Path(cfg.aptos_root) / \"test_images\" / f\"{name}.png\"))\n    return frame[[\"image_id\", \"image_path\"]]\n\n\ndef build_training_frame(cfg: Config) -> pd.DataFrame:\n    frame = build_aptos_train_frame(cfg)\n    frame = frame.sample(frac=1.0, random_state=cfg.seed).reset_index(drop=True)\n    return frame\n","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:51.301009Z","iopub.execute_input":"2026-04-09T08:50:51.301226Z","iopub.status.idle":"2026-04-09T08:50:51.318502Z","shell.execute_reply.started":"2026-04-09T08:50:51.301204Z","shell.execute_reply":"2026-04-09T08:50:51.317547Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"267dd6e1","cell_type":"code","source":"train_frame = build_training_frame(CFG)\naptos_test_frame = build_aptos_test_frame(CFG)\n\nprint(\"Training samples:\", len(train_frame))\nprint(\"APTOS test samples:\", len(aptos_test_frame))\ndataset_source_distribution = train_frame.groupby([\"source\", \"label\"]).size().unstack(fill_value=0)\ndataset_source_distribution = dataset_source_distribution.rename(\n    columns={class_idx: CLASS_NAMES[class_idx] for class_idx in range(CFG.num_classes)}\n)\ndataset_overall_class_counts = (\n    train_frame[\"label\"]\n    .value_counts()\n    .sort_index()\n    .rename(\"count\")\n    .rename_axis(\"label\")\n    .reset_index()\n)\ndataset_overall_class_counts[\"class_name\"] = dataset_overall_class_counts[\"label\"].map(\n    lambda class_idx: CLASS_NAMES[int(class_idx)]\n)\ndataset_overall_class_counts = dataset_overall_class_counts[[\"label\", \"class_name\", \"count\"]]\ndataset_source_class_counts = dataset_source_distribution.reset_index()\n\ndisplay(train_frame.head())\ndisplay(dataset_source_distribution)\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 4))\nsns.countplot(\n    data=train_frame,\n    x=\"label\",\n    hue=\"label\",\n    order=range(CFG.num_classes),\n    hue_order=range(CFG.num_classes),\n    ax=axes[0],\n    palette=\"Blues_d\",\n    legend=False,\n)\naxes[0].set_title(\"Class Distribution\")\naxes[0].set_xlabel(\"DR Grade\")\naxes[0].set_ylabel(\"Images\")\n\nsns.countplot(data=train_frame, x=\"source\", hue=\"label\", ax=axes[1], palette=\"viridis\")\naxes[1].set_title(\"Dataset Sources\")\naxes[1].set_xlabel(\"Source\")\naxes[1].set_ylabel(\"Images\")\nplt.tight_layout()\nsave_figure(fig, \"dataset_distribution.png\")\nplt.show()\n\nexport_table_bundle(dataset_overall_class_counts, \"dataset_overall_class_counts\", include_index=False)\nexport_table_bundle(dataset_source_class_counts, \"dataset_source_class_counts\", include_index=False)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:51.320451Z","iopub.execute_input":"2026-04-09T08:50:51.320682Z","iopub.status.idle":"2026-04-09T08:50:55.228633Z","shell.execute_reply.started":"2026-04-09T08:50:51.320660Z","shell.execute_reply":"2026-04-09T08:50:55.227920Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"083d6396","cell_type":"code","source":"def choose_stratify_targets(frame: pd.DataFrame) -> pd.Series:\n    if \"source\" in frame.columns and frame[\"source\"].nunique() > 1:\n        joint_targets = frame[\"source\"].astype(str) + \"__\" + frame[\"label\"].astype(str)\n        if joint_targets.value_counts().min() >= 2:\n            return joint_targets\n    return frame[\"label\"].astype(str)\n\n\ndef stratified_split_frame(frame: pd.DataFrame, test_size: float, seed: int) -> tuple[pd.DataFrame, pd.DataFrame]:\n    splitter = StratifiedShuffleSplit(\n        n_splits=1,\n        test_size=test_size,\n        random_state=seed,\n    )\n    stratify_targets = choose_stratify_targets(frame)\n    train_idx, test_idx = next(splitter.split(frame, stratify_targets))\n    train_df = frame.iloc[train_idx].reset_index(drop=True)\n    test_df = frame.iloc[test_idx].reset_index(drop=True)\n    return train_df, test_df\n\n\ndef make_research_splits(\n    frame: pd.DataFrame,\n    val_size: float,\n    calibration_size: float,\n    seed: int,\n) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    if val_size <= 0 or calibration_size < 0 or (val_size + calibration_size) >= 1:\n        raise ValueError(\"Expected val_size > 0, calibration_size >= 0, and val_size + calibration_size < 1.\")\n\n    train_calibration_df, val_df = stratified_split_frame(frame, test_size=val_size, seed=seed)\n\n    if calibration_size == 0:\n        calibration_df = pd.DataFrame(columns=frame.columns)\n        train_df = train_calibration_df.reset_index(drop=True)\n    else:\n        calibration_fraction = calibration_size / (1.0 - val_size)\n        train_df, calibration_df = stratified_split_frame(\n            train_calibration_df,\n            test_size=calibration_fraction,\n            seed=seed + 1,\n        )\n\n    return train_df, calibration_df.reset_index(drop=True), val_df.reset_index(drop=True)\n\n\ntrain_df, calibration_df, val_df = make_research_splits(\n    train_frame,\n    val_size=CFG.val_size,\n    calibration_size=CFG.calibration_size,\n    seed=CFG.seed,\n)\n\nprint(\"Train split:\", train_df.shape)\nprint(\"Calibration split:\", calibration_df.shape)\nprint(\"Validation split:\", val_df.shape)\nsplit_class_distribution = pd.concat(\n    [\n        train_df[\"label\"].value_counts().sort_index().rename(\"train\"),\n        calibration_df[\"label\"].value_counts().sort_index().rename(\"calibration\"),\n        val_df[\"label\"].value_counts().sort_index().rename(\"validation\"),\n    ],\n    axis=1,\n).fillna(0).astype(int)\nsplit_class_distribution.index = [CLASS_NAMES[int(class_idx)] for class_idx in split_class_distribution.index]\nsplit_class_distribution.index.name = \"class_name\"\n\nsplit_source_label_distribution = pd.concat(\n    [\n        train_df.groupby([\"source\", \"label\"]).size().rename(\"train\"),\n        calibration_df.groupby([\"source\", \"label\"]).size().rename(\"calibration\"),\n        val_df.groupby([\"source\", \"label\"]).size().rename(\"validation\"),\n    ],\n    axis=1,\n).fillna(0).astype(int).reset_index()\nsplit_source_label_distribution[\"class_name\"] = split_source_label_distribution[\"label\"].map(\n    lambda class_idx: CLASS_NAMES[int(class_idx)]\n)\nsplit_source_label_distribution = split_source_label_distribution[\n    [\"source\", \"label\", \"class_name\", \"train\", \"calibration\", \"validation\"]\n]\n\ndisplay(split_class_distribution)\ndisplay(split_source_label_distribution)\nexport_table_bundle(split_class_distribution, \"split_class_distribution\", include_index=True)\nexport_table_bundle(split_source_label_distribution, \"split_source_label_distribution\", include_index=False)\n\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(\n        (CFG.image_size, CFG.image_size),\n        scale=(0.92, 1.0),\n        ratio=(0.96, 1.04),\n    ),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomApply(\n        [\n            transforms.RandomAffine(\n                degrees=10,\n                translate=(0.04, 0.04),\n                scale=(0.95, 1.05),\n            )\n        ],\n        p=0.35,\n    ),\n    transforms.ColorJitter(brightness=0.16, contrast=0.16, saturation=0.12, hue=0.015),\n    transforms.RandomAutocontrast(p=0.08),\n    transforms.RandomApply([transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 0.6))], p=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    transforms.RandomErasing(p=0.08, scale=(0.02, 0.04), ratio=(0.5, 2.0), value=\"random\"),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize((CFG.image_size, CFG.image_size)),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:55.229723Z","iopub.execute_input":"2026-04-09T08:50:55.230546Z","iopub.status.idle":"2026-04-09T08:50:55.288514Z","shell.execute_reply.started":"2026-04-09T08:50:55.230514Z","shell.execute_reply":"2026-04-09T08:50:55.287691Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"70d3f5e7","cell_type":"code","source":"class RetinopathyDataset(Dataset):\n    def __init__(self, frame: pd.DataFrame, transform, cfg: Config, cache_dir: Path | None = None):\n        self.frame = frame.reset_index(drop=True).copy()\n        self.transform = transform\n        self.cfg = cfg\n        self.cache_dir = Path(cache_dir) if cache_dir is not None else None\n        if self.cache_dir is not None:\n            self.cache_dir.mkdir(parents=True, exist_ok=True)\n\n    def __len__(self) -> int:\n        return len(self.frame)\n\n    def _cache_path(self, row: pd.Series) -> Path | None:\n        if self.cache_dir is None:\n            return None\n        image_stem = Path(str(row[\"image_path\"])).stem\n        cache_name = f\"{image_stem}_{self.cfg.image_size}_clahe{int(self.cfg.apply_clahe)}.npy\"\n        return self.cache_dir / cache_name\n\n    def _load_image(self, row: pd.Series) -> Image.Image:\n        cache_path = self._cache_path(row)\n        if cache_path is not None and cache_path.exists():\n            with cache_path.open(\"rb\") as handle:\n                cached_image = np.load(handle, allow_pickle=False)\n            return Image.fromarray(cached_image)\n\n        image = load_retina_image(row[\"image_path\"], self.cfg.image_size, apply_clahe=self.cfg.apply_clahe)\n        if cache_path is not None:\n            image_array = np.asarray(image, dtype=np.uint8)\n            tmp_path = cache_path.with_suffix(f\"{cache_path.suffix}.{os.getpid()}.tmp\")\n            try:\n                with tmp_path.open(\"wb\") as handle:\n                    np.save(handle, image_array, allow_pickle=False)\n                os.replace(tmp_path, cache_path)\n            except OSError:\n                if tmp_path.exists():\n                    tmp_path.unlink(missing_ok=True)\n            return Image.fromarray(image_array)\n        return image\n\n    def __getitem__(self, index: int) -> dict:\n        row = self.frame.iloc[index]\n        image = self._load_image(row)\n        image_tensor = self.transform(image)\n\n        batch = {\n            \"image\": image_tensor,\n            \"image_id\": row[\"image_id\"],\n            \"image_path\": row[\"image_path\"],\n        }\n        if \"label\" in row.index:\n            batch[\"label\"] = torch.tensor(int(row[\"label\"]), dtype=torch.long)\n        return batch\n\n\ndef build_sampling_tools(\n    labels: pd.Series,\n    num_classes: int,\n    balance_power: float = 1.0,\n    max_weight: float | None = None,\n) -> tuple[WeightedRandomSampler, torch.Tensor]:\n    counts = labels.value_counts().reindex(range(num_classes), fill_value=0)\n    adjusted_counts = counts.replace(0, 1)\n    inverse_frequency = len(labels) / (num_classes * adjusted_counts)\n    class_weights = inverse_frequency.pow(balance_power)\n    class_weights = class_weights / class_weights.mean()\n    if max_weight is not None:\n        class_weights = class_weights.clip(lower=0.75, upper=float(max_weight))\n        class_weights = class_weights / class_weights.mean()\n    sample_weights = labels.map(class_weights).to_numpy()\n    sampler = WeightedRandomSampler(\n        weights=torch.as_tensor(sample_weights, dtype=torch.double),\n        num_samples=len(sample_weights),\n        replacement=True,\n    )\n    return sampler, torch.as_tensor(class_weights.to_numpy(), dtype=torch.float32)\n\n\ndef build_loss_class_weights(\n    labels: pd.Series,\n    num_classes: int,\n    balance_power: float = 0.35,\n    min_weight: float = 0.85,\n    max_weight: float = 1.75,\n    advanced_class_start: int = 3,\n    advanced_boost: float = 0.0,\n) -> torch.Tensor:\n    counts = labels.value_counts().reindex(range(num_classes), fill_value=0)\n    adjusted_counts = counts.replace(0, 1)\n    inverse_frequency = len(labels) / (num_classes * adjusted_counts)\n    class_weights = inverse_frequency.pow(balance_power)\n    class_weights = class_weights / class_weights.mean()\n    if advanced_boost > 0:\n        advanced_mask = class_weights.index.to_series().ge(int(advanced_class_start))\n        class_weights.loc[advanced_mask] = class_weights.loc[advanced_mask] * (1.0 + float(advanced_boost))\n    class_weights = class_weights.clip(lower=float(min_weight), upper=float(max_weight))\n    class_weights = class_weights / class_weights.mean()\n    return torch.as_tensor(class_weights.to_numpy(), dtype=torch.float32)\n\n\npreprocessed_cache_root = WORKDIR / \"preprocessed_cache\" if CFG.cache_preprocessed_images else None\ntrain_dataset = RetinopathyDataset(\n    train_df,\n    train_transform,\n    CFG,\n    cache_dir=preprocessed_cache_root / \"train\" if preprocessed_cache_root is not None else None,\n)\ncalibration_dataset = RetinopathyDataset(\n    calibration_df,\n    eval_transform,\n    CFG,\n    cache_dir=preprocessed_cache_root / \"calibration\" if preprocessed_cache_root is not None else None,\n)\nval_dataset = RetinopathyDataset(\n    val_df,\n    eval_transform,\n    CFG,\n    cache_dir=preprocessed_cache_root / \"validation\" if preprocessed_cache_root is not None else None,\n)\ntrain_sampler, class_weights = build_sampling_tools(\n    train_df[\"label\"],\n    CFG.num_classes,\n    balance_power=CFG.class_balance_power,\n    max_weight=CFG.sampler_weight_cap,\n)\nloss_class_weights = build_loss_class_weights(\n    train_df[\"label\"],\n    CFG.num_classes,\n    balance_power=CFG.loss_class_balance_power,\n    min_weight=CFG.loss_weight_floor,\n    max_weight=CFG.loss_weight_cap,\n    advanced_class_start=CFG.advanced_class_start,\n    advanced_boost=CFG.advanced_class_boost,\n)\nloader_kwargs = {\n    \"num_workers\": CFG.num_workers,\n    \"pin_memory\": torch.cuda.is_available(),\n    \"persistent_workers\": CFG.num_workers > 0,\n}\nif CFG.num_workers > 0:\n    loader_kwargs[\"prefetch_factor\"] = max(2, int(CFG.prefetch_factor))\n\ntrain_sampler = train_sampler if CFG.use_weighted_sampler else None\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.batch_size,\n    sampler=train_sampler,\n    shuffle=train_sampler is None,\n    drop_last=True,\n    **loader_kwargs,\n)\n\ncalibration_loader = DataLoader(\n    calibration_dataset,\n    batch_size=CFG.batch_size,\n    shuffle=False,\n    **loader_kwargs,\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=CFG.batch_size,\n    shuffle=False,\n    **loader_kwargs,\n)\n\nprint(\"Sampler class weights:\", class_weights.tolist())\nprint(\"Loss class weights:\", loss_class_weights.tolist())\nprint(\"Weighted sampler enabled:\", CFG.use_weighted_sampler)\nprint(\"DataLoader workers per loader:\", loader_kwargs[\"num_workers\"])\nprint(\"Persistent workers enabled:\", loader_kwargs[\"persistent_workers\"])\nprint(\"Prefetch factor:\", loader_kwargs.get(\"prefetch_factor\", \"n/a\"))\nprint(\"Preprocessed image cache enabled:\", CFG.cache_preprocessed_images)\nprint(\"Train batches:\", len(train_loader))\nprint(\"Calibration batches:\", len(calibration_loader))\nprint(\"Validation batches:\", len(val_loader))","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:55.290200Z","iopub.execute_input":"2026-04-09T08:50:55.290560Z","iopub.status.idle":"2026-04-09T08:50:55.355123Z","shell.execute_reply.started":"2026-04-09T08:50:55.290534Z","shell.execute_reply":"2026-04-09T08:50:55.354534Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"01f7edec","cell_type":"markdown","source":"## Hybrid Model\n\nThe architecture uses:\n\n- **EfficientNet-B3** by default to capture stronger local lesion-level texture and vascular cues\n- **DeiT-Tiny** to capture wider retinal context\n- a learned fusion block that mixes local, global, interaction, and disagreement signals\n- a standard **5-class classifier head**\n- a parallel **ordinal head** that regularizes the training process for ordered DR grades\n\nRecommended backbone choices for a free Kaggle GPU:\n\n- `efficientnet_b3`: best balance of accuracy and runtime for this notebook\n- `tf_efficientnetv2_s`: stronger but slower, better if you stay on APTOS only\n- `mobilevitv2_150`: lighter fallback if memory or time gets tight","metadata":{}},{"id":"fd4de96e","cell_type":"code","source":"def create_timm_backbone(model_name: str, pretrained: bool = True):\n    try:\n        return timm.create_model(model_name, pretrained=pretrained)\n    except Exception as error:\n        print(f\"Falling back to pretrained=False for {model_name}: {error}\")\n        return timm.create_model(model_name, pretrained=False)\n\n\nclass HybridDRClassifier(nn.Module):\n    def __init__(\n        self,\n        num_classes: int = 5,\n        pretrained: bool = True,\n        fusion_dim: int = 384,\n        cnn_backbone: str = \"efficientnet_b3\",\n        vit_backbone: str = \"deit_tiny_patch16_224\",\n    ):\n        super().__init__()\n        self.cnn_name = cnn_backbone\n        self.vit_name = vit_backbone\n        self.cnn = create_timm_backbone(cnn_backbone, pretrained=pretrained)\n        self.vit = create_timm_backbone(vit_backbone, pretrained=pretrained)\n\n        cnn_dim = self.cnn.num_features\n        vit_dim = self.vit.num_features\n\n        self.cnn_projection = nn.Sequential(\n            nn.Linear(cnn_dim, fusion_dim),\n            nn.LayerNorm(fusion_dim),\n            nn.GELU(),\n            nn.Dropout(0.20),\n        )\n\n        self.vit_projection = nn.Sequential(\n            nn.Linear(vit_dim, fusion_dim),\n            nn.LayerNorm(fusion_dim),\n            nn.GELU(),\n            nn.Dropout(0.20),\n        )\n\n        fusion_input_dim = fusion_dim * 4\n        self.fusion_block = nn.Sequential(\n            nn.Linear(fusion_input_dim, fusion_dim * 2),\n            nn.GELU(),\n            nn.Dropout(0.25),\n            nn.Linear(fusion_dim * 2, fusion_dim),\n            nn.GELU(),\n            nn.Dropout(0.20),\n        )\n\n        self.classifier = nn.Linear(fusion_dim, num_classes)\n        self.ordinal_head = nn.Linear(fusion_dim, num_classes - 1)\n\n    def forward(self, images: torch.Tensor) -> dict[str, torch.Tensor]:\n        cnn_maps = self.cnn.forward_features(images)\n        cnn_vector = self.cnn.forward_head(cnn_maps, pre_logits=True)\n        if cnn_vector.ndim > 2:\n            cnn_vector = torch.flatten(cnn_vector, 1)\n\n        vit_tokens = self.vit.forward_features(images)\n        vit_cls = vit_tokens[:, 0]\n\n        cnn_repr = self.cnn_projection(cnn_vector)\n        vit_repr = self.vit_projection(vit_cls)\n\n        interaction = cnn_repr * vit_repr\n        disagreement = torch.abs(cnn_repr - vit_repr)\n        fused_features = self.fusion_block(torch.cat([cnn_repr, vit_repr, interaction, disagreement], dim=1))\n\n        logits = self.classifier(fused_features)\n        ordinal_logits = self.ordinal_head(fused_features)\n\n        return {\n            \"logits\": logits,\n            \"ordinal_logits\": ordinal_logits,\n            \"cnn_maps\": cnn_maps,\n            \"vit_tokens\": vit_tokens,\n            \"fused_features\": fused_features,\n        }\n\n\nmodel = HybridDRClassifier(\n    num_classes=CFG.num_classes,\n    pretrained=CFG.pretrained_backbones,\n    cnn_backbone=CFG.cnn_backbone,\n    vit_backbone=CFG.vit_backbone,\n).to(DEVICE)\nif USE_CHANNELS_LAST:\n    model = model.to(memory_format=torch.channels_last)\ntotal_params = sum(parameter.numel() for parameter in model.parameters()) / 1_000_000\nprint(f\"Trainable parameters: {total_params:.2f}M\")\nprint(\"CNN backbone:\", CFG.cnn_backbone)\nprint(\"Transformer backbone:\", CFG.vit_backbone)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:55.355986Z","iopub.execute_input":"2026-04-09T08:50:55.356425Z","iopub.status.idle":"2026-04-09T08:50:58.200480Z","shell.execute_reply.started":"2026-04-09T08:50:55.356385Z","shell.execute_reply":"2026-04-09T08:50:58.199620Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"49053a52","cell_type":"code","source":"def make_ordinal_targets(labels: torch.Tensor, num_classes: int) -> torch.Tensor:\n    thresholds = torch.arange(num_classes - 1, device=labels.device).unsqueeze(0)\n    return (labels.unsqueeze(1) > thresholds).float()\n\n\nclass HybridOrdinalLoss(nn.Module):\n    def __init__(\n        self,\n        num_classes: int,\n        class_weights: torch.Tensor | None = None,\n        ordinal_weight: float = 0.40,\n        distance_weight: float = 0.15,\n        hard_example_gamma: float = 1.25,\n        label_smoothing: float = 0.0,\n        advanced_class_start: int = 3,\n        advanced_class_boost: float = 0.0,\n        undergrading_penalty_weight: float = 0.0,\n        undergrading_margin_power: float = 1.5,\n    ):\n        super().__init__()\n        self.num_classes = num_classes\n        self.ordinal_weight = ordinal_weight\n        self.distance_weight = distance_weight\n        self.hard_example_gamma = hard_example_gamma\n        self.label_smoothing = label_smoothing\n        self.advanced_class_start = int(advanced_class_start)\n        self.advanced_class_boost = float(advanced_class_boost)\n        self.undergrading_penalty_weight = float(undergrading_penalty_weight)\n        self.undergrading_margin_power = float(undergrading_margin_power)\n        self.register_buffer(\"grade_values\", torch.arange(num_classes, dtype=torch.float32))\n\n        if class_weights is None:\n            self.class_weights = None\n        else:\n            self.register_buffer(\"class_weights\", class_weights.float())\n\n    def forward(self, outputs: dict[str, torch.Tensor], labels: torch.Tensor) -> torch.Tensor:\n        logits = outputs[\"logits\"]\n        ordinal_logits = outputs[\"ordinal_logits\"]\n\n        ce_loss = F.cross_entropy(\n            logits,\n            labels,\n            reduction=\"none\",\n            weight=self.class_weights,\n            label_smoothing=self.label_smoothing,\n        )\n        ordinal_targets = make_ordinal_targets(labels, self.num_classes)\n        ordinal_loss = F.binary_cross_entropy_with_logits(ordinal_logits, ordinal_targets, reduction=\"none\").mean(dim=1)\n\n        probabilities = logits.softmax(dim=1)\n        true_class_prob = probabilities.gather(1, labels.unsqueeze(1)).squeeze(1)\n        grade_values = self.grade_values.to(device=probabilities.device, dtype=probabilities.dtype)\n        expected_grade = (probabilities * grade_values).sum(dim=1)\n        distance_penalty = (expected_grade - labels.float()).pow(2)\n        undergrading_margin = F.relu(labels.float() - expected_grade)\n\n        adjacent_mass = torch.zeros_like(true_class_prob)\n        for shift in (-1, 1):\n            neighbor = labels + shift\n            valid = (neighbor >= 0) & (neighbor < self.num_classes)\n            adjacent_mass[valid] += probabilities[valid, neighbor[valid]]\n\n        advanced_mask = (labels >= self.advanced_class_start).float()\n        advanced_focus = 1.0 + advanced_mask * self.advanced_class_boost\n        hard_weight = (\n            1.0\n            + self.hard_example_gamma * (1.0 - true_class_prob).pow(2)\n            + 0.30 * adjacent_mass\n            + 0.20 * distance_penalty.sqrt()\n        )\n        main_loss = ce_loss + self.ordinal_weight * ordinal_loss + self.distance_weight * distance_penalty\n        undergrading_penalty = advanced_mask * undergrading_margin.pow(self.undergrading_margin_power)\n        total_loss = advanced_focus * hard_weight * main_loss + self.undergrading_penalty_weight * undergrading_penalty\n        return total_loss.mean()\n\n\ndef safe_multiclass_auc(y_true: np.ndarray, y_prob: np.ndarray, num_classes: int) -> float:\n    try:\n        y_true_binary = label_binarize(y_true, classes=list(range(num_classes)))\n        return float(roc_auc_score(y_true_binary, y_prob, average=\"macro\", multi_class=\"ovr\"))\n    except ValueError:\n        return float(\"nan\")\n\n\ndef negative_log_likelihood(y_true: np.ndarray, y_prob: np.ndarray) -> float:\n    if len(y_true) == 0:\n        return float(\"nan\")\n    clipped = np.clip(y_prob, 1e-7, 1.0)\n    return float(-np.log(clipped[np.arange(len(y_true)), y_true]).mean())\n\n\ndef expected_calibration_error(y_true: np.ndarray, y_pred: np.ndarray, confidences: np.ndarray, num_bins: int = 10) -> float:\n    if len(y_true) == 0:\n        return float(\"nan\")\n\n    confidences = np.asarray(confidences, dtype=np.float32)\n    correctness = (y_pred == y_true).astype(np.float32)\n    bin_edges = np.linspace(0.0, 1.0, num_bins + 1)\n    ece = 0.0\n\n    for start, end in zip(bin_edges[:-1], bin_edges[1:]):\n        if end == 1.0:\n            in_bin = (confidences >= start) & (confidences <= end)\n        else:\n            in_bin = (confidences >= start) & (confidences < end)\n        if not np.any(in_bin):\n            continue\n        bin_accuracy = correctness[in_bin].mean()\n        bin_confidence = confidences[in_bin].mean()\n        ece += np.abs(bin_accuracy - bin_confidence) * in_bin.mean()\n\n    return float(ece)\n\n\ndef logits_to_probabilities(logits: np.ndarray, temperature: float = 1.0) -> np.ndarray:\n    temperature = max(float(temperature), 1e-3)\n    logits_tensor = torch.as_tensor(logits, dtype=torch.float32)\n    return torch.softmax(logits_tensor / temperature, dim=1).cpu().numpy()\n\n\ndef compute_metrics(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    y_prob: np.ndarray,\n    num_classes: int,\n    predicted_confidences: np.ndarray | None = None,\n) -> dict[str, float]:\n    if len(y_true) == 0:\n        return {\n            \"accuracy\": float(\"nan\"),\n            \"balanced_accuracy\": float(\"nan\"),\n            \"macro_f1\": float(\"nan\"),\n            \"weighted_f1\": float(\"nan\"),\n            \"qwk\": float(\"nan\"),\n            \"macro_auc_ovr\": float(\"nan\"),\n            \"adjacent_error_rate\": float(\"nan\"),\n            \"minority_recall\": float(\"nan\"),\n            \"min_class_recall\": float(\"nan\"),\n            \"nll\": float(\"nan\"),\n            \"ece\": float(\"nan\"),\n        }\n\n    if predicted_confidences is None:\n        predicted_confidences = y_prob.max(axis=1)\n    else:\n        predicted_confidences = np.asarray(predicted_confidences, dtype=np.float32)\n\n    _, recalls, _, _ = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=list(range(num_classes)),\n        zero_division=0,\n    )\n    recalls = np.asarray(recalls, dtype=np.float32)\n    if len(recalls) >= 2:\n        minority_recall = float(recalls[-2:].mean())\n    else:\n        minority_recall = float(recalls.mean())\n\n    adjacent_errors = ((np.abs(y_true - y_pred) == 1) & (y_true != y_pred)).mean()\n    return {\n        \"accuracy\": float(accuracy_score(y_true, y_pred)),\n        \"balanced_accuracy\": float(balanced_accuracy_score(y_true, y_pred)),\n        \"macro_f1\": float(f1_score(y_true, y_pred, average=\"macro\", zero_division=0)),\n        \"weighted_f1\": float(f1_score(y_true, y_pred, average=\"weighted\", zero_division=0)),\n        \"qwk\": float(cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")),\n        \"macro_auc_ovr\": safe_multiclass_auc(y_true, y_prob, num_classes),\n        \"adjacent_error_rate\": float(adjacent_errors),\n        \"minority_recall\": minority_recall,\n        \"min_class_recall\": float(recalls.min()),\n        \"nll\": negative_log_likelihood(y_true, y_prob),\n        \"ece\": expected_calibration_error(y_true, y_pred, predicted_confidences),\n    }\n","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:58.201573Z","iopub.execute_input":"2026-04-09T08:50:58.201909Z","iopub.status.idle":"2026-04-09T08:50:58.224606Z","shell.execute_reply.started":"2026-04-09T08:50:58.201868Z","shell.execute_reply":"2026-04-09T08:50:58.223726Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"fc466051","cell_type":"code","source":"class TrainingProgressWidget:\n    def __init__(self, total_epochs: int, enabled: bool = True):\n        self.enabled = bool(enabled)\n        self.total_epochs = max(int(total_epochs), 1)\n        self.current_epoch = 0\n        self.total_steps = 1\n        self.phase_name = \"Idle\"\n        self.update_interval = 1\n\n        if not self.enabled:\n            return\n\n        try:\n            self.header = widgets.HTML(\"<b>Hybrid DR training dashboard</b>\")\n            self.status = widgets.HTML(\"Waiting to start\")\n            self.epoch_bar = widgets.IntProgress(\n                value=0,\n                min=0,\n                max=self.total_epochs,\n                description=\"Epoch\",\n                bar_style=\"info\",\n            )\n            self.step_bar = widgets.IntProgress(\n                value=0,\n                min=0,\n                max=1,\n                description=\"Batches\",\n                bar_style=\"\",\n            )\n            self.metrics = widgets.HTML(\"No metrics yet\")\n            display(widgets.VBox([self.header, self.status, self.epoch_bar, self.step_bar, self.metrics]))\n        except Exception as error:\n            self.enabled = False\n            print(f\"Falling back to text logs because the notebook widget could not initialize: {error}\")\n\n    def start_epoch(self, epoch: int) -> None:\n        self.current_epoch = int(epoch)\n        if not self.enabled:\n            return\n        self.epoch_bar.value = max(0, epoch - 1)\n        self.status.value = f\"<b>Epoch {epoch}/{self.total_epochs}</b> | preparing...\"\n\n    def start_phase(self, phase_name: str, total_steps: int) -> None:\n        self.phase_name = phase_name\n        self.total_steps = max(int(total_steps), 1)\n        self.update_interval = 1 if self.total_steps <= 20 else max(1, self.total_steps // 20)\n        if not self.enabled:\n            return\n        self.step_bar.max = self.total_steps\n        self.step_bar.value = 0\n        self.step_bar.description = phase_name\n        self.step_bar.bar_style = \"info\" if phase_name.lower().startswith(\"train\") else \"warning\"\n        self.status.value = f\"<b>Epoch {self.current_epoch}/{self.total_epochs}</b> | {phase_name} started\"\n\n    def update_batch(self, step: int, loss: float, lr: float | None = None) -> None:\n        if not self.enabled:\n            return\n        if step % self.update_interval != 0 and step != self.total_steps:\n            return\n        self.step_bar.value = min(step, self.total_steps)\n        lr_text = f\" | lr={lr:.6f}\" if lr is not None else \"\"\n        self.status.value = (\n            f\"<b>Epoch {self.current_epoch}/{self.total_epochs}</b> | \"\n            f\"{self.phase_name} {step}/{self.total_steps} | loss={loss:.4f}{lr_text}\"\n        )\n\n    def update_epoch_summary(self, row: dict[str, float], best_score: float, best_accuracy: float) -> None:\n        if not self.enabled:\n            return\n        self.epoch_bar.value = min(int(row[\"epoch\"]), self.total_epochs)\n        self.step_bar.value = self.step_bar.max\n        self.step_bar.bar_style = \"success\"\n        self.metrics.value = (\n            \"<table>\"\n            f\"<tr><td><b>Stage</b></td><td>{row['stage']}</td></tr>\"\n            f\"<tr><td><b>Train loss</b></td><td>{row['train_loss']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val loss</b></td><td>{row['val_loss']:.4f}</td></tr>\"\n            f\"<tr><td><b>Train acc</b></td><td>{row['train_accuracy']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val acc</b></td><td>{row['val_accuracy']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val bal acc</b></td><td>{row['val_balanced_accuracy']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val QWK</b></td><td>{row['val_qwk']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val macro F1</b></td><td>{row['val_macro_f1']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val weighted F1</b></td><td>{row['val_weighted_f1']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val minority recall</b></td><td>{row['val_minority_recall']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val min recall</b></td><td>{row['val_min_class_recall']:.4f}</td></tr>\"\n            f\"<tr><td><b>Acc gap</b></td><td>{row['accuracy_gap']:.4f}</td></tr>\"\n            f\"<tr><td><b>Val ECE</b></td><td>{row['val_ece']:.4f}</td></tr>\"\n            f\"<tr><td><b>Selection score</b></td><td>{row['selection_score']:.4f}</td></tr>\"\n            f\"<tr><td><b>Head lr</b></td><td>{row['head_lr']:.6f}</td></tr>\"\n            f\"<tr><td><b>Backbone lr</b></td><td>{row['backbone_lr']:.6f}</td></tr>\"\n            f\"<tr><td><b>Best score</b></td><td>{best_score:.4f}</td></tr>\"\n            f\"<tr><td><b>Best val acc</b></td><td>{best_accuracy:.4f}</td></tr>\"\n            \"</table>\"\n        )\n        self.status.value = f\"<b>Epoch {self.current_epoch}/{self.total_epochs}</b> | finished\"\n\n    def finish(self) -> None:\n        if not self.enabled:\n            return\n        self.step_bar.bar_style = \"success\"\n        self.status.value = self.status.value + \" | training complete\"\n\n\ndef compute_selection_score(row: dict[str, float]) -> float:\n    generalization_penalty = (\n        0.22 * max(row[\"train_accuracy\"] - row[\"val_accuracy\"], 0.0)\n        + 0.10 * max(row[\"train_balanced_accuracy\"] - row[\"val_balanced_accuracy\"], 0.0)\n        + 0.18 * max(row[\"train_qwk\"] - row[\"val_qwk\"], 0.0)\n        + 0.12 * max(row[\"train_macro_f1\"] - row[\"val_macro_f1\"], 0.0)\n        + 0.22 * max(row[\"train_minority_recall\"] - row[\"val_minority_recall\"], 0.0)\n        + 0.16 * max(row[\"train_min_class_recall\"] - row[\"val_min_class_recall\"], 0.0)\n    )\n    calibration_penalty = max(row[\"val_ece\"] - 0.08, 0.0)\n    stage_penalty = 0.04 if row.get(\"stage\") == \"head_warmup\" else 0.0\n    return float(\n        0.28 * row[\"val_qwk\"]\n        + 0.16 * row[\"val_accuracy\"]\n        + 0.12 * row[\"val_macro_f1\"]\n        + 0.12 * row[\"val_balanced_accuracy\"]\n        + 0.22 * row[\"val_minority_recall\"]\n        + 0.10 * row[\"val_min_class_recall\"]\n        - 0.14 * generalization_penalty\n        - 0.03 * calibration_penalty\n        - stage_penalty\n    )\n\n\ndef split_optimizer_groups(model: nn.Module) -> tuple[list[nn.Parameter], list[nn.Parameter]]:\n    head_prefixes = (\n        \"cnn_projection\",\n        \"vit_projection\",\n        \"fusion_block\",\n        \"classifier\",\n        \"ordinal_head\",\n    )\n    backbone_params, head_params = [], []\n    for name, parameter in model.named_parameters():\n        if not parameter.requires_grad:\n            continue\n        if any(name.startswith(prefix) for prefix in head_prefixes):\n            head_params.append(parameter)\n        else:\n            backbone_params.append(parameter)\n    return backbone_params, head_params\n\n\ndef set_backbone_trainable(model: nn.Module, enabled: bool) -> None:\n    for parameter in model.cnn.parameters():\n        parameter.requires_grad = enabled\n    for parameter in model.vit.parameters():\n        parameter.requires_grad = enabled\n\n    for block in [model.cnn_projection, model.vit_projection, model.fusion_block, model.classifier, model.ordinal_head]:\n        for parameter in block.parameters():\n            parameter.requires_grad = True\n\n\ndef get_current_lrs(optimizer: torch.optim.Optimizer) -> tuple[float, float]:\n    if not optimizer.param_groups:\n        return float(\"nan\"), float(\"nan\")\n    backbone_lr = float(optimizer.param_groups[0][\"lr\"])\n    head_lr = float(optimizer.param_groups[-1][\"lr\"])\n    return backbone_lr, head_lr\n\n\nclass ExponentialMovingAverage:\n    def __init__(self, model: nn.Module, decay: float):\n        self.decay = float(decay)\n        self.module = copy.deepcopy(model).eval()\n        for parameter in self.module.parameters():\n            parameter.requires_grad_(False)\n\n    @torch.no_grad()\n    def update(self, model: nn.Module) -> None:\n        model_state = model.state_dict()\n        for name, ema_value in self.module.state_dict().items():\n            model_value = model_state[name].detach()\n            if ema_value.dtype.is_floating_point:\n                ema_value.mul_(self.decay).add_(model_value, alpha=1.0 - self.decay)\n            else:\n                ema_value.copy_(model_value)\n\n\ndef build_epoch_scheduler(optimizer: torch.optim.Optimizer, cfg: Config):\n    if cfg.epochs <= 1:\n        return None\n\n    warmup_epochs = min(cfg.warmup_epochs, max(cfg.epochs - 1, 0))\n    if warmup_epochs > 0:\n        warmup = torch.optim.lr_scheduler.LinearLR(\n            optimizer,\n            start_factor=0.35,\n            end_factor=1.0,\n            total_iters=warmup_epochs,\n        )\n        cosine = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer,\n            T_max=max(cfg.epochs - warmup_epochs, 1),\n            eta_min=cfg.min_lr,\n        )\n        return torch.optim.lr_scheduler.SequentialLR(\n            optimizer,\n            schedulers=[warmup, cosine],\n            milestones=[warmup_epochs],\n        )\n\n    return torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=max(cfg.epochs, 1),\n        eta_min=cfg.min_lr,\n    )\n\n\ndef prepare_batch_images(images: torch.Tensor) -> torch.Tensor:\n    images = images.to(DEVICE, non_blocking=torch.cuda.is_available())\n    if USE_CHANNELS_LAST:\n        images = images.contiguous(memory_format=torch.channels_last)\n    return images\n\n\ndef get_autocast_context():\n    return torch.amp.autocast(device_type=DEVICE.type, enabled=USE_AMP) if USE_AMP else nullcontext()\n\n\ndef run_epoch(\n    model: nn.Module,\n    loader: DataLoader,\n    criterion: nn.Module,\n    optimizer: torch.optim.Optimizer | None = None,\n    scaler: torch.amp.GradScaler | None = None,\n    max_grad_norm: float | None = None,\n    progress_widget: TrainingProgressWidget | None = None,\n    phase_name: str = \"Train\",\n    ema_model: ExponentialMovingAverage | None = None,\n) -> dict:\n    training = optimizer is not None\n    model.train(training)\n    if training and hasattr(model, \"cnn\") and not any(parameter.requires_grad for parameter in model.cnn.parameters()):\n        model.cnn.eval()\n    if training and hasattr(model, \"vit\") and not any(parameter.requires_grad for parameter in model.vit.parameters()):\n        model.vit.eval()\n\n    all_targets = []\n    all_predictions = []\n    all_probabilities = []\n    running_loss = []\n\n    if progress_widget is not None:\n        progress_widget.start_phase(phase_name, len(loader))\n\n    for step, batch in enumerate(loader, start=1):\n        images = prepare_batch_images(batch[\"image\"])\n        labels = batch[\"label\"].to(DEVICE, non_blocking=True)\n\n        if training:\n            optimizer.zero_grad(set_to_none=True)\n\n        grad_context = nullcontext() if training else torch.no_grad()\n        autocast_context = get_autocast_context()\n\n        with grad_context:\n            with autocast_context:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n            if training:\n                scaler.scale(loss).backward()\n                if max_grad_norm is not None and max_grad_norm > 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)\n                scaler.step(optimizer)\n                scaler.update()\n                if ema_model is not None:\n                    ema_model.update(model)\n\n        probabilities = outputs[\"logits\"].float().softmax(dim=1).detach().cpu().numpy()\n        predictions = probabilities.argmax(axis=1)\n\n        all_targets.append(labels.detach().cpu().numpy())\n        all_predictions.append(predictions)\n        all_probabilities.append(probabilities)\n        running_loss.append(float(loss.item()))\n\n        if progress_widget is not None:\n            current_lr = float(optimizer.param_groups[0][\"lr\"]) if training else None\n            progress_widget.update_batch(step, float(np.mean(running_loss)), current_lr)\n\n    y_true = np.concatenate(all_targets)\n    y_pred = np.concatenate(all_predictions)\n    y_prob = np.concatenate(all_probabilities)\n    metrics = compute_metrics(y_true, y_pred, y_prob, CFG.num_classes)\n\n    return {\n        \"loss\": float(np.mean(running_loss)),\n        \"metrics\": metrics,\n        \"y_true\": y_true,\n        \"y_pred\": y_pred,\n        \"y_prob\": y_prob,\n    }\n\n\ndef fit_model(\n    model: nn.Module,\n    train_loader: DataLoader,\n    val_loader: DataLoader,\n    criterion: nn.Module,\n    optimizer: torch.optim.Optimizer,\n    scheduler,\n    cfg: Config,\n) -> tuple[pd.DataFrame, Path]:\n    scaler = torch.amp.GradScaler(DEVICE.type, enabled=USE_AMP)\n    best_qwk = -float(\"inf\")\n    best_accuracy = -float(\"inf\")\n    best_score = -float(\"inf\")\n    history = []\n    checkpoint_path = WORKDIR / \"best_hybrid_dr_classifier.pt\"\n    epochs_without_improvement = 0\n    progress_widget = TrainingProgressWidget(cfg.epochs, enabled=cfg.use_training_widgets)\n    ema_model = ExponentialMovingAverage(model, cfg.ema_decay) if cfg.use_ema else None\n    current_stage = None\n\n    for epoch in range(1, cfg.epochs + 1):\n        progress_widget.start_epoch(epoch)\n        backbone_trainable = epoch > cfg.freeze_backbone_epochs\n        stage_name = \"full_finetune\" if backbone_trainable else \"head_warmup\"\n        if stage_name != current_stage:\n            set_backbone_trainable(model, enabled=backbone_trainable)\n            current_stage = stage_name\n            print(f\"Training stage: {stage_name} (epoch {epoch})\")\n\n        train_result = run_epoch(\n            model,\n            train_loader,\n            criterion,\n            optimizer=optimizer,\n            scaler=scaler,\n            max_grad_norm=cfg.max_grad_norm,\n            progress_widget=progress_widget,\n            phase_name=\"Train\",\n            ema_model=ema_model,\n        )\n        eval_model = ema_model.module if ema_model is not None else model\n        val_result = run_epoch(\n            eval_model,\n            val_loader,\n            criterion,\n            optimizer=None,\n            scaler=scaler,\n            progress_widget=progress_widget,\n            phase_name=\"Validation\",\n        )\n        backbone_lr, head_lr = get_current_lrs(optimizer)\n\n        row = {\n            \"epoch\": epoch,\n            \"stage\": stage_name,\n            \"train_loss\": train_result[\"loss\"],\n            \"val_loss\": val_result[\"loss\"],\n            \"train_accuracy\": train_result[\"metrics\"][\"accuracy\"],\n            \"train_balanced_accuracy\": train_result[\"metrics\"][\"balanced_accuracy\"],\n            \"train_qwk\": train_result[\"metrics\"][\"qwk\"],\n            \"val_qwk\": val_result[\"metrics\"][\"qwk\"],\n            \"train_macro_f1\": train_result[\"metrics\"][\"macro_f1\"],\n            \"train_weighted_f1\": train_result[\"metrics\"][\"weighted_f1\"],\n            \"train_minority_recall\": train_result[\"metrics\"][\"minority_recall\"],\n            \"train_min_class_recall\": train_result[\"metrics\"][\"min_class_recall\"],\n            \"val_macro_f1\": val_result[\"metrics\"][\"macro_f1\"],\n            \"val_weighted_f1\": val_result[\"metrics\"][\"weighted_f1\"],\n            \"val_minority_recall\": val_result[\"metrics\"][\"minority_recall\"],\n            \"val_min_class_recall\": val_result[\"metrics\"][\"min_class_recall\"],\n            \"val_accuracy\": val_result[\"metrics\"][\"accuracy\"],\n            \"val_balanced_accuracy\": val_result[\"metrics\"][\"balanced_accuracy\"],\n            \"val_auc\": val_result[\"metrics\"][\"macro_auc_ovr\"],\n            \"train_ece\": train_result[\"metrics\"][\"ece\"],\n            \"val_ece\": val_result[\"metrics\"][\"ece\"],\n            \"train_nll\": train_result[\"metrics\"][\"nll\"],\n            \"val_nll\": val_result[\"metrics\"][\"nll\"],\n            \"accuracy_gap\": train_result[\"metrics\"][\"accuracy\"] - val_result[\"metrics\"][\"accuracy\"],\n            \"qwk_gap\": train_result[\"metrics\"][\"qwk\"] - val_result[\"metrics\"][\"qwk\"],\n            \"macro_f1_gap\": train_result[\"metrics\"][\"macro_f1\"] - val_result[\"metrics\"][\"macro_f1\"],\n            \"backbone_lr\": backbone_lr,\n            \"head_lr\": head_lr,\n            \"monitor_model\": \"ema\" if ema_model is not None else \"base\",\n        }\n        row[\"selection_score\"] = compute_selection_score(row)\n        history.append(row)\n\n        print(\n            f\"Epoch {epoch:02d} | \"\n            f\"stage={row['stage']} | \"\n            f\"train_acc={row['train_accuracy']:.4f} | val_acc={row['val_accuracy']:.4f} | \"\n            f\"val_bal_acc={row['val_balanced_accuracy']:.4f} | val_qwk={row['val_qwk']:.4f} | \"\n            f\"val_macro_f1={row['val_macro_f1']:.4f} | val_minority_recall={row['val_minority_recall']:.4f} | \"\n            f\"gap={row['accuracy_gap']:.4f} | score={row['selection_score']:.4f} | head_lr={row['head_lr']:.6f} | \"\n            f\"backbone_lr={row['backbone_lr']:.6f}\"\n        )\n\n        best_qwk = max(best_qwk, row[\"val_qwk\"])\n        best_accuracy = max(best_accuracy, row[\"val_accuracy\"])\n\n        if row[\"selection_score\"] > best_score:\n            best_score = row[\"selection_score\"]\n            epochs_without_improvement = 0\n            torch.save(\n                {\n                    \"model_state\": eval_model.state_dict(),\n                    \"config\": asdict(cfg),\n                    \"history_row\": row,\n                    \"best_score\": best_score,\n                    \"best_qwk\": best_qwk,\n                    \"best_val_accuracy\": best_accuracy,\n                },\n                checkpoint_path,\n            )\n        else:\n            epochs_without_improvement += 1\n\n        progress_widget.update_epoch_summary(row, best_score, best_accuracy)\n\n        if scheduler is not None:\n            scheduler.step()\n\n        early_stop_ready = epoch >= max(cfg.min_epochs_before_early_stop, cfg.freeze_backbone_epochs + 2)\n        if early_stop_ready and epochs_without_improvement >= cfg.early_stopping_patience:\n            print(f\"Early stopping triggered after {cfg.early_stopping_patience} non-improving epochs.\")\n            break\n\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    progress_widget.finish()\n    return pd.DataFrame(history), checkpoint_path","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:58.225917Z","iopub.execute_input":"2026-04-09T08:50:58.226251Z","iopub.status.idle":"2026-04-09T08:50:58.271027Z","shell.execute_reply.started":"2026-04-09T08:50:58.226226Z","shell.execute_reply":"2026-04-09T08:50:58.270146Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e5701a16","cell_type":"code","source":"criterion = HybridOrdinalLoss(\n    num_classes=CFG.num_classes,\n    class_weights=loss_class_weights.to(DEVICE) if CFG.use_loss_class_weights else None,\n    ordinal_weight=CFG.ordinal_weight,\n    distance_weight=CFG.distance_weight,\n    hard_example_gamma=CFG.hard_example_gamma,\n    label_smoothing=CFG.label_smoothing,\n    advanced_class_start=CFG.advanced_class_start,\n    advanced_class_boost=CFG.advanced_class_boost,\n    undergrading_penalty_weight=CFG.undergrading_penalty_weight,\n    undergrading_margin_power=CFG.undergrading_margin_power,\n).to(DEVICE)\nbackbone_params, head_params = split_optimizer_groups(model)\noptimizer = torch.optim.AdamW(\n    [\n        {\n            \"params\": backbone_params,\n            \"lr\": CFG.backbone_lr,\n            \"weight_decay\": CFG.weight_decay,\n        },\n        {\n            \"params\": head_params,\n            \"lr\": CFG.head_lr,\n            \"weight_decay\": CFG.head_weight_decay,\n        },\n    ],\n    betas=(0.9, 0.99),\n)\nscheduler = build_epoch_scheduler(optimizer, CFG)\n\nhistory_df, checkpoint_path = fit_model(\n    model=model,\n    train_loader=train_loader,\n    val_loader=val_loader,\n    criterion=criterion,\n    optimizer=optimizer,\n    scheduler=scheduler,\n    cfg=CFG,\n)\n\ndisplay(history_df)\ndisplay(history_df.sort_values(\"selection_score\", ascending=False).head(1).round(4))\ndisplay(history_df.tail(1).round(4))\nprint(\"Best checkpoint:\", checkpoint_path)\nprint(\"Loss class weights enabled:\", CFG.use_loss_class_weights)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T08:50:58.271917Z","iopub.execute_input":"2026-04-09T08:50:58.272295Z","iopub.status.idle":"2026-04-09T09:06:18.193343Z","shell.execute_reply.started":"2026-04-09T08:50:58.272252Z","shell.execute_reply":"2026-04-09T09:06:18.192504Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"3568a2da","cell_type":"code","source":"def predict_logits_with_tta(\n    model: nn.Module,\n    images: torch.Tensor,\n    tta_horizontal: bool = False,\n) -> torch.Tensor:\n    logits = model(images)[\"logits\"]\n    if tta_horizontal:\n        flipped_images = torch.flip(images, dims=[3])\n        if USE_CHANNELS_LAST:\n            flipped_images = flipped_images.contiguous(memory_format=torch.channels_last)\n        logits = 0.5 * (logits + model(flipped_images)[\"logits\"])\n    return logits.float()\n\n\ndef collect_labeled_logits(\n    model: nn.Module,\n    loader: DataLoader,\n    tta_horizontal: bool = False,\n) -> tuple[pd.DataFrame, np.ndarray]:\n    model.eval()\n    all_rows = []\n    all_logits = []\n\n    with torch.inference_mode():\n        for batch in loader:\n            images = prepare_batch_images(batch[\"image\"])\n            labels = batch[\"label\"].cpu().numpy()\n            with get_autocast_context():\n                logits = predict_logits_with_tta(model, images, tta_horizontal=tta_horizontal).cpu().numpy()\n\n            for image_id, image_path, label in zip(\n                batch[\"image_id\"],\n                batch[\"image_path\"],\n                labels,\n            ):\n                all_rows.append(\n                    {\n                        \"image_id\": image_id,\n                        \"image_path\": image_path,\n                        \"true_label\": int(label),\n                    }\n                )\n            all_logits.append(logits)\n\n    if not all_logits:\n        return pd.DataFrame(columns=[\"image_id\", \"image_path\", \"true_label\"]), np.empty((0, CFG.num_classes), dtype=np.float32)\n    return pd.DataFrame(all_rows), np.concatenate(all_logits)\n\n\ndef expected_grade_from_probabilities(probability_array: np.ndarray) -> np.ndarray:\n    if len(probability_array) == 0:\n        return np.empty((0,), dtype=np.float32)\n    grade_values = np.arange(probability_array.shape[1], dtype=np.float32)\n    return (probability_array * grade_values[None, :]).sum(axis=1).astype(np.float32)\n\n\ndef predict_from_expected_grade(expected_grade: np.ndarray, decision_thresholds: list[float] | np.ndarray) -> np.ndarray:\n    thresholds = np.asarray(decision_thresholds, dtype=np.float32)\n    return np.digitize(expected_grade, bins=thresholds, right=False).astype(int)\n\n\ndef threshold_selection_score(metrics: dict[str, float]) -> float:\n    return float(\n        0.24 * metrics[\"qwk\"]\n        + 0.14 * metrics[\"accuracy\"]\n        + 0.14 * metrics[\"macro_f1\"]\n        + 0.12 * metrics[\"balanced_accuracy\"]\n        + 0.22 * metrics[\"minority_recall\"]\n        + 0.10 * metrics[\"min_class_recall\"]\n        - 0.04 * metrics[\"ece\"]\n    )\n\n\ndef fit_ordinal_thresholds(\n    probabilities: np.ndarray,\n    y_true: np.ndarray,\n    cfg: Config,\n) -> tuple[np.ndarray, pd.DataFrame]:\n    default_thresholds = np.arange(cfg.num_classes - 1, dtype=np.float32) + 0.5\n    if len(y_true) == 0:\n        threshold_frame = pd.DataFrame(\n            {\n                \"boundary\": [f\"{CLASS_NAMES[i]} -> {CLASS_NAMES[i + 1]}\" for i in range(cfg.num_classes - 1)],\n                \"default_threshold\": default_thresholds,\n                \"optimized_threshold\": default_thresholds,\n                \"searched\": False,\n            }\n        )\n        return default_thresholds, threshold_frame\n\n    expected_grade = expected_grade_from_probabilities(probabilities)\n    sorted_scores = np.unique(np.sort(expected_grade))\n    midpoint_candidates = (\n        (sorted_scores[:-1] + sorted_scores[1:]) / 2.0 if len(sorted_scores) > 1 else np.array([], dtype=np.float32)\n    )\n    global_candidate_pool = np.unique(\n        np.concatenate(\n            [\n                np.linspace(-0.5, cfg.num_classes - 0.5, cfg.threshold_search_steps, dtype=np.float32),\n                midpoint_candidates,\n                default_thresholds,\n            ]\n        )\n    )\n\n    thresholds = default_thresholds.copy()\n    predictions = predict_from_expected_grade(expected_grade, thresholds)\n    best_metrics = compute_metrics(y_true, predictions, probabilities, cfg.num_classes)\n    best_score = threshold_selection_score(best_metrics)\n    epsilon = 1e-3\n    fixed_count = int(cfg.threshold_fixed_lower_boundaries) if cfg.threshold_focus_upper_boundaries_only else 0\n    optimized_mask = np.zeros(len(thresholds), dtype=bool)\n\n    for _ in range(cfg.threshold_search_passes):\n        improved = False\n        for threshold_idx in range(len(thresholds)):\n            if threshold_idx < fixed_count:\n                continue\n\n            left_boundary = thresholds[threshold_idx - 1] + epsilon if threshold_idx > 0 else -0.5\n            right_boundary = (\n                thresholds[threshold_idx + 1] - epsilon if threshold_idx + 1 < len(thresholds) else cfg.num_classes - 0.5\n            )\n            default_center = default_thresholds[threshold_idx]\n            candidate_pool = global_candidate_pool[\n                (global_candidate_pool >= default_center - cfg.threshold_search_radius)\n                & (global_candidate_pool <= default_center + cfg.threshold_search_radius)\n            ]\n            valid_candidates = candidate_pool[(candidate_pool > left_boundary) & (candidate_pool < right_boundary)]\n            if len(valid_candidates) == 0:\n                continue\n\n            local_best_threshold = thresholds[threshold_idx]\n            local_best_metrics = best_metrics\n            local_best_score = best_score\n            for candidate in valid_candidates:\n                trial_thresholds = thresholds.copy()\n                trial_thresholds[threshold_idx] = float(candidate)\n                trial_predictions = predict_from_expected_grade(expected_grade, trial_thresholds)\n                trial_metrics = compute_metrics(y_true, trial_predictions, probabilities, cfg.num_classes)\n                trial_score = threshold_selection_score(trial_metrics)\n                if trial_score > local_best_score + 1e-6:\n                    local_best_threshold = float(candidate)\n                    local_best_metrics = trial_metrics\n                    local_best_score = trial_score\n\n            if abs(local_best_threshold - thresholds[threshold_idx]) > epsilon:\n                thresholds[threshold_idx] = local_best_threshold\n                best_metrics = local_best_metrics\n                best_score = local_best_score\n                optimized_mask[threshold_idx] = True\n                improved = True\n\n        if not improved:\n            break\n\n    threshold_frame = pd.DataFrame(\n        {\n            \"boundary\": [f\"{CLASS_NAMES[i]} -> {CLASS_NAMES[i + 1]}\" for i in range(cfg.num_classes - 1)],\n            \"default_threshold\": default_thresholds,\n            \"optimized_threshold\": thresholds,\n            \"searched\": optimized_mask,\n        }\n    )\n    return thresholds, threshold_frame\n\n\ndef select_decision_thresholds(\n    probabilities: np.ndarray,\n    y_true: np.ndarray,\n    cfg: Config,\n) -> tuple[np.ndarray | None, pd.DataFrame, pd.DataFrame]:\n    default_thresholds = np.arange(cfg.num_classes - 1, dtype=np.float32) + 0.5\n    threshold_frame = pd.DataFrame(\n        {\n            \"boundary\": [f\"{CLASS_NAMES[i]} -> {CLASS_NAMES[i + 1]}\" for i in range(cfg.num_classes - 1)],\n            \"default_threshold\": default_thresholds,\n            \"optimized_threshold\": default_thresholds,\n            \"selected_threshold\": default_thresholds,\n            \"selected_policy\": \"argmax\",\n            \"selection_reason\": \"threshold optimization disabled\",\n        }\n    )\n\n    if len(y_true) == 0:\n        comparison_frame = pd.DataFrame([{\"policy\": \"argmax\", \"selected\": True}])\n        return None, threshold_frame, comparison_frame\n\n    argmax_predictions = probabilities.argmax(axis=1).astype(int)\n    argmax_metrics = compute_metrics(y_true, argmax_predictions, probabilities, cfg.num_classes)\n    comparison_rows = [{\"policy\": \"argmax\", \"selected\": True, **argmax_metrics}]\n\n    if not cfg.enable_threshold_optimization:\n        return None, threshold_frame, pd.DataFrame(comparison_rows)\n\n    optimized_thresholds, optimized_frame = fit_ordinal_thresholds(probabilities, y_true, cfg)\n    expected_grade = expected_grade_from_probabilities(probabilities)\n    optimized_predictions = predict_from_expected_grade(expected_grade, optimized_thresholds)\n    optimized_confidences = probabilities[np.arange(len(optimized_predictions)), optimized_predictions].astype(np.float32)\n    optimized_metrics = compute_metrics(\n        y_true,\n        optimized_predictions,\n        probabilities,\n        cfg.num_classes,\n        predicted_confidences=optimized_confidences,\n    )\n\n    max_shift = float(np.abs(optimized_thresholds - default_thresholds).max())\n    min_gap = float(np.min(np.diff(optimized_thresholds))) if len(optimized_thresholds) > 1 else float(\"inf\")\n    accuracy_drop = float(argmax_metrics[\"accuracy\"] - optimized_metrics[\"accuracy\"])\n    qwk_drop = float(argmax_metrics[\"qwk\"] - optimized_metrics[\"qwk\"])\n    ece_increase = float(optimized_metrics[\"ece\"] - argmax_metrics[\"ece\"])\n    score_gain = float(threshold_selection_score(optimized_metrics) - threshold_selection_score(argmax_metrics))\n    advanced_recall_gain = float(optimized_metrics[\"minority_recall\"] - argmax_metrics[\"minority_recall\"])\n    min_recall_gain = float(optimized_metrics[\"min_class_recall\"] - argmax_metrics[\"min_class_recall\"])\n    searched_boundary_count = int(optimized_frame.get(\"searched\", pd.Series(dtype=bool)).sum()) if \"searched\" in optimized_frame.columns else 0\n    use_optimized = (\n        score_gain >= cfg.threshold_min_score_gain\n        and advanced_recall_gain >= cfg.threshold_min_advanced_recall_gain\n        and min_recall_gain >= cfg.threshold_min_min_recall_gain\n        and accuracy_drop <= cfg.threshold_max_accuracy_drop\n        and qwk_drop <= cfg.threshold_max_qwk_drop\n        and ece_increase <= cfg.threshold_max_ece_increase\n        and max_shift <= cfg.threshold_max_shift\n        and min_gap >= cfg.threshold_min_gap\n        and searched_boundary_count > 0\n    )\n\n    selected_policy = \"calibrated_thresholds\" if use_optimized else \"argmax\"\n    selection_reason = (\n        f\"accepted | score_gain={score_gain:.4f} | advanced_recall_gain={advanced_recall_gain:.4f} | min_recall_gain={min_recall_gain:.4f}\"\n        if use_optimized\n        else (\n            f\"rejected | score_gain={score_gain:.4f} | advanced_recall_gain={advanced_recall_gain:.4f} | \"\n            f\"min_recall_gain={min_recall_gain:.4f} | accuracy_drop={accuracy_drop:.4f} | qwk_drop={qwk_drop:.4f} | ece_increase={ece_increase:.4f}\"\n        )\n    )\n\n    threshold_frame = optimized_frame.copy()\n    threshold_frame[\"selected_threshold\"] = optimized_thresholds if use_optimized else default_thresholds\n    threshold_frame[\"selected_policy\"] = selected_policy\n    threshold_frame[\"selection_reason\"] = selection_reason\n    threshold_frame[\"max_shift\"] = max_shift\n    threshold_frame[\"min_gap\"] = min_gap\n    threshold_frame[\"accuracy_drop_vs_argmax\"] = accuracy_drop\n    threshold_frame[\"qwk_drop_vs_argmax\"] = qwk_drop\n    threshold_frame[\"ece_increase_vs_argmax\"] = ece_increase\n    threshold_frame[\"score_gain_vs_argmax\"] = score_gain\n    threshold_frame[\"advanced_recall_gain_vs_argmax\"] = advanced_recall_gain\n    threshold_frame[\"min_class_recall_gain_vs_argmax\"] = min_recall_gain\n\n    comparison_rows.append({\"policy\": \"calibrated_thresholds\", \"selected\": use_optimized, **optimized_metrics})\n    return optimized_thresholds if use_optimized else None, threshold_frame, pd.DataFrame(comparison_rows)\n\n\ndef build_labeled_prediction_frame(\n    metadata_frame: pd.DataFrame,\n    logits: np.ndarray,\n    temperature: float = 1.0,\n    decision_thresholds: list[float] | np.ndarray | None = None,\n) -> tuple[pd.DataFrame, np.ndarray, dict[str, float]]:\n    if len(metadata_frame) == 0 or len(logits) == 0:\n        empty_frame = metadata_frame.copy()\n        empty_frame[\"pred_label\"] = pd.Series(dtype=int)\n        empty_frame[\"confidence\"] = pd.Series(dtype=float)\n        empty_frame[\"expected_grade\"] = pd.Series(dtype=float)\n        empty_frame[\"decision_policy\"] = pd.Series(dtype=str)\n        empty_frame[\"distance\"] = pd.Series(dtype=int)\n        empty_probabilities = np.empty((0, CFG.num_classes), dtype=np.float32)\n        empty_metrics = compute_metrics(np.array([], dtype=int), np.array([], dtype=int), empty_probabilities, CFG.num_classes)\n        return empty_frame, empty_probabilities, empty_metrics\n\n    probability_array = logits_to_probabilities(logits, temperature=temperature)\n    expected_grade = expected_grade_from_probabilities(probability_array)\n    if decision_thresholds is None:\n        predicted_labels = probability_array.argmax(axis=1).astype(int)\n        decision_policy = \"argmax\"\n    else:\n        predicted_labels = predict_from_expected_grade(expected_grade, decision_thresholds)\n        decision_policy = \"calibrated_thresholds\"\n\n    predicted_confidences = probability_array[np.arange(len(predicted_labels)), predicted_labels].astype(np.float32)\n    prediction_frame = metadata_frame.copy()\n    prediction_frame[\"pred_label\"] = predicted_labels\n    prediction_frame[\"confidence\"] = predicted_confidences.astype(float)\n    prediction_frame[\"expected_grade\"] = expected_grade.astype(float)\n    prediction_frame[\"decision_policy\"] = decision_policy\n    prediction_frame[\"distance\"] = np.abs(\n        prediction_frame[\"pred_label\"].to_numpy() - prediction_frame[\"true_label\"].to_numpy()\n    ).astype(int)\n    metrics = compute_metrics(\n        prediction_frame[\"true_label\"].to_numpy(),\n        prediction_frame[\"pred_label\"].to_numpy(),\n        probability_array,\n        CFG.num_classes,\n        predicted_confidences=predicted_confidences,\n    )\n    return prediction_frame, probability_array, metrics\n\n\ndef build_classification_report_frame(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    class_names: list[str],\n) -> pd.DataFrame:\n    report = classification_report(\n        y_true,\n        y_pred,\n        labels=list(range(len(class_names))),\n        target_names=class_names,\n        output_dict=True,\n        zero_division=0,\n    )\n    return pd.DataFrame(report).T\n\n\ndef build_confusion_frames(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    class_names: list[str],\n) -> tuple[pd.DataFrame, pd.DataFrame]:\n    matrix = confusion_matrix(y_true, y_pred, labels=list(range(len(class_names))))\n    count_frame = pd.DataFrame(matrix, index=class_names, columns=class_names)\n    normalized_frame = count_frame.div(count_frame.sum(axis=1).replace(0, 1), axis=0)\n    return count_frame, normalized_frame\n\n\ndef fit_temperature_scaler(\n    calibration_logits: np.ndarray,\n    calibration_labels: np.ndarray,\n    cfg: Config,\n) -> float:\n    if len(calibration_labels) == 0:\n        return 1.0\n\n    logits_tensor = torch.tensor(calibration_logits, dtype=torch.float32)\n    labels_tensor = torch.tensor(calibration_labels, dtype=torch.long)\n    temperature = nn.Parameter(torch.ones(1, dtype=torch.float32))\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.LBFGS([temperature], lr=0.05, max_iter=50, line_search_fn=\"strong_wolfe\")\n\n    def closure():\n        optimizer.zero_grad()\n        clamped_temperature = temperature.clamp(cfg.temperature_min, cfg.temperature_max)\n        loss = criterion(logits_tensor / clamped_temperature, labels_tensor)\n        loss.backward()\n        return loss\n\n    optimizer.step(closure)\n    return float(temperature.detach().clamp(cfg.temperature_min, cfg.temperature_max).item())\n\n\ndef summarize_probability_calibration(\n    logits: np.ndarray,\n    y_true: np.ndarray,\n    temperature: float,\n    split_name: str,\n) -> list[dict[str, float | str]]:\n    rows = []\n    for setting_name, used_temperature in [\n        (\"uncalibrated\", 1.0),\n        (\"temperature_scaled\", temperature),\n    ]:\n        probabilities = logits_to_probabilities(logits, temperature=used_temperature)\n        predictions = probabilities.argmax(axis=1)\n        metrics = compute_metrics(y_true, predictions, probabilities, CFG.num_classes)\n        rows.append(\n            {\n                \"split\": split_name,\n                \"setting\": setting_name,\n                \"temperature\": float(used_temperature),\n                \"accuracy\": metrics[\"accuracy\"],\n                \"balanced_accuracy\": metrics[\"balanced_accuracy\"],\n                \"qwk\": metrics[\"qwk\"],\n                \"macro_f1\": metrics[\"macro_f1\"],\n                \"weighted_f1\": metrics[\"weighted_f1\"],\n                \"macro_auc_ovr\": metrics[\"macro_auc_ovr\"],\n                \"nll\": metrics[\"nll\"],\n                \"ece\": metrics[\"ece\"],\n            }\n        )\n    return rows\n\n\ndef plot_training_history(history: pd.DataFrame):\n    fig, axes = plt.subplots(2, 2, figsize=(16, 10))\n    axes = axes.ravel()\n\n    axes[0].plot(history[\"epoch\"], history[\"train_loss\"], marker=\"o\", label=\"Train loss\")\n    axes[0].plot(history[\"epoch\"], history[\"val_loss\"], marker=\"o\", label=\"Val loss\")\n    axes[0].set_title(\"Loss\")\n    axes[0].set_xlabel(\"Epoch\")\n    axes[0].set_ylabel(\"Loss\")\n    axes[0].legend()\n\n    axes[1].plot(history[\"epoch\"], history[\"train_accuracy\"], marker=\"o\", label=\"Train acc\")\n    axes[1].plot(history[\"epoch\"], history[\"val_accuracy\"], marker=\"o\", label=\"Val acc\")\n    axes[1].plot(history[\"epoch\"], history[\"val_balanced_accuracy\"], marker=\"o\", label=\"Val balanced acc\")\n    axes[1].set_title(\"Accuracy\")\n    axes[1].set_xlabel(\"Epoch\")\n    axes[1].set_ylabel(\"Score\")\n    axes[1].legend()\n\n    axes[2].plot(history[\"epoch\"], history[\"val_qwk\"], marker=\"o\", label=\"Val QWK\")\n    axes[2].plot(history[\"epoch\"], history[\"val_macro_f1\"], marker=\"o\", label=\"Val macro F1\")\n    axes[2].plot(history[\"epoch\"], history[\"val_weighted_f1\"], marker=\"o\", label=\"Val weighted F1\")\n    axes[2].plot(history[\"epoch\"], history[\"selection_score\"], marker=\"o\", label=\"Selection score\")\n    axes[2].set_title(\"Validation Quality\")\n    axes[2].set_xlabel(\"Epoch\")\n    axes[2].set_ylabel(\"Score\")\n    axes[2].legend()\n\n    axes[3].plot(history[\"epoch\"], history[\"accuracy_gap\"], marker=\"o\", label=\"Accuracy gap\")\n    axes[3].plot(history[\"epoch\"], history[\"qwk_gap\"], marker=\"o\", label=\"QWK gap\")\n    axes[3].plot(history[\"epoch\"], history[\"macro_f1_gap\"], marker=\"o\", label=\"Macro F1 gap\")\n    axes[3].set_title(\"Generalization Gap\")\n    axes[3].set_xlabel(\"Epoch\")\n    axes[3].set_ylabel(\"Gap\")\n    axes[3].legend()\n\n    plt.tight_layout()\n    plt.show()\n    return fig\n\n\ndef plot_confusion_heatmap(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    class_names: list[str],\n) -> tuple[object, pd.DataFrame, pd.DataFrame]:\n    count_frame, normalized_frame = build_confusion_frames(y_true, y_pred, class_names)\n    fig, axes = plt.subplots(1, 2, figsize=(15, 6))\n    sns.heatmap(count_frame, annot=True, fmt=\"d\", cmap=\"mako\", ax=axes[0])\n    axes[0].set_title(\"Validation Confusion Matrix (Counts)\")\n    axes[0].set_xlabel(\"Predicted Grade\")\n    axes[0].set_ylabel(\"True Grade\")\n\n    sns.heatmap(normalized_frame, annot=True, fmt=\".2f\", cmap=\"mako\", vmin=0.0, vmax=1.0, ax=axes[1])\n    axes[1].set_title(\"Validation Confusion Matrix (Row-Normalized)\")\n    axes[1].set_xlabel(\"Predicted Grade\")\n    axes[1].set_ylabel(\"True Grade\")\n    plt.tight_layout()\n    plt.show()\n    return fig, count_frame, normalized_frame\n\n\ndef plot_multiclass_roc(y_true: np.ndarray, y_prob: np.ndarray, class_names: list[str]):\n    y_true_binary = label_binarize(y_true, classes=list(range(len(class_names))))\n    fig, ax = plt.subplots(figsize=(8, 6))\n    plotted_any = False\n\n    for class_idx, class_name in enumerate(class_names):\n        positives = y_true_binary[:, class_idx].sum()\n        if positives == 0 or positives == len(y_true_binary):\n            continue\n        fpr, tpr, _ = roc_curve(y_true_binary[:, class_idx], y_prob[:, class_idx])\n        auc_value = roc_auc_score(y_true_binary[:, class_idx], y_prob[:, class_idx])\n        ax.plot(fpr, tpr, linewidth=2, label=f\"{class_name} (AUC={auc_value:.3f})\")\n        plotted_any = True\n\n    ax.plot([0, 1], [0, 1], linestyle=\"--\", color=\"gray\", linewidth=1.5, label=\"Chance\")\n    ax.set_title(\"Validation One-vs-Rest ROC Curves\")\n    ax.set_xlabel(\"False Positive Rate\")\n    ax.set_ylabel(\"True Positive Rate\")\n    if plotted_any:\n        ax.legend(loc=\"lower right\")\n    plt.tight_layout()\n    plt.show()\n    return fig\n\n\ndef plot_reliability_diagram(\n    y_true: np.ndarray,\n    y_prob: np.ndarray,\n    num_bins: int = 10,\n) -> tuple[object, pd.DataFrame]:\n    predictions = y_prob.argmax(axis=1)\n    confidences = y_prob.max(axis=1)\n    correctness = (predictions == y_true).astype(float)\n    bin_edges = np.linspace(0.0, 1.0, num_bins + 1)\n    rows = []\n\n    for start, end in zip(bin_edges[:-1], bin_edges[1:]):\n        if end == 1.0:\n            in_bin = (confidences >= start) & (confidences <= end)\n        else:\n            in_bin = (confidences >= start) & (confidences < end)\n\n        if not np.any(in_bin):\n            continue\n\n        rows.append(\n            {\n                \"bin_start\": float(start),\n                \"bin_end\": float(end),\n                \"bin_center\": float((start + end) / 2.0),\n                \"mean_confidence\": float(confidences[in_bin].mean()),\n                \"empirical_accuracy\": float(correctness[in_bin].mean()),\n                \"count\": int(in_bin.sum()),\n            }\n        )\n\n    reliability_frame = pd.DataFrame(rows)\n    fig, ax = plt.subplots(figsize=(7, 6))\n    ax.plot([0, 1], [0, 1], linestyle=\"--\", color=\"gray\", linewidth=1.5, label=\"Perfect calibration\")\n    if len(reliability_frame):\n        ax.plot(\n            reliability_frame[\"mean_confidence\"],\n            reliability_frame[\"empirical_accuracy\"],\n            marker=\"o\",\n            linewidth=2,\n            label=\"Observed\",\n        )\n    ax.set_title(\"Validation Reliability Diagram\")\n    ax.set_xlabel(\"Mean Confidence\")\n    ax.set_ylabel(\"Empirical Accuracy\")\n    ax.set_xlim(0, 1)\n    ax.set_ylim(0, 1)\n    ax.legend(loc=\"upper left\")\n    plt.tight_layout()\n    plt.show()\n    return fig, reliability_frame\n\n\ncheckpoint = torch.load(checkpoint_path, map_location=DEVICE)\nmodel.load_state_dict(checkpoint[\"model_state\"])\n\ncalibration_metadata_frame, calibration_logits = collect_labeled_logits(\n    model,\n    calibration_loader,\n    tta_horizontal=CFG.tta_horizontal,\n)\ncalibration_labels = calibration_metadata_frame[\"true_label\"].to_numpy()\ntemperature = fit_temperature_scaler(calibration_logits, calibration_labels, CFG)\ncalibration_argmax_frame, calibration_probabilities, calibration_argmax_metrics = build_labeled_prediction_frame(\n    calibration_metadata_frame,\n    calibration_logits,\n    temperature=temperature,\n)\ndecision_thresholds, decision_threshold_frame, calibration_policy_comparison = select_decision_thresholds(\n    calibration_probabilities,\n    calibration_labels,\n    CFG,\n)\ncandidate_thresholds = (\n    decision_threshold_frame[\"optimized_threshold\"].to_numpy(dtype=np.float32)\n    if \"optimized_threshold\" in decision_threshold_frame.columns\n    else None\n)\ncalibration_candidate_frame = calibration_argmax_frame\ncalibration_candidate_metrics = calibration_argmax_metrics\nif candidate_thresholds is not None and CFG.enable_threshold_optimization:\n    calibration_candidate_frame, calibration_probabilities, calibration_candidate_metrics = build_labeled_prediction_frame(\n        calibration_metadata_frame,\n        calibration_logits,\n        temperature=temperature,\n        decision_thresholds=candidate_thresholds,\n    )\n\nselected_decision_policy = \"calibrated_thresholds\" if decision_thresholds is not None else \"argmax\"\nif selected_decision_policy == \"calibrated_thresholds\":\n    calibration_prediction_frame = calibration_candidate_frame\n    calibration_metrics = calibration_candidate_metrics\nelse:\n    calibration_prediction_frame = calibration_argmax_frame\n    calibration_metrics = calibration_argmax_metrics\n\nval_metadata_frame, val_logits = collect_labeled_logits(\n    model,\n    val_loader,\n    tta_horizontal=CFG.tta_horizontal,\n)\nval_argmax_frame, val_probabilities, val_argmax_metrics = build_labeled_prediction_frame(\n    val_metadata_frame,\n    val_logits,\n    temperature=temperature,\n)\nval_candidate_frame = val_argmax_frame\nval_candidate_metrics = val_argmax_metrics\nif candidate_thresholds is not None and CFG.enable_threshold_optimization:\n    val_candidate_frame, val_probabilities, val_candidate_metrics = build_labeled_prediction_frame(\n        val_metadata_frame,\n        val_logits,\n        temperature=temperature,\n        decision_thresholds=candidate_thresholds,\n    )\n\nif selected_decision_policy == \"calibrated_thresholds\":\n    val_prediction_frame = val_candidate_frame\n    val_metrics = val_candidate_metrics\nelse:\n    val_prediction_frame = val_argmax_frame\n    val_metrics = val_argmax_metrics\n\nclass_report_frame = build_classification_report_frame(\n    val_prediction_frame[\"true_label\"].to_numpy(),\n    val_prediction_frame[\"pred_label\"].to_numpy(),\n    CLASS_NAMES,\n)\n\ncalibration_report = pd.DataFrame(\n    summarize_probability_calibration(calibration_logits, calibration_labels, temperature, \"Calibration\")\n    + summarize_probability_calibration(\n        val_logits,\n        val_prediction_frame[\"true_label\"].to_numpy(),\n        temperature,\n        \"Validation\",\n    )\n)\ndecision_policy_rows = [\n    {\"split\": \"Calibration\", \"policy\": \"argmax\", \"selected\": selected_decision_policy == \"argmax\", **calibration_argmax_metrics},\n    {\"split\": \"Validation\", \"policy\": \"argmax\", \"selected\": selected_decision_policy == \"argmax\", **val_argmax_metrics},\n]\nif candidate_thresholds is not None and CFG.enable_threshold_optimization:\n    decision_policy_rows.extend(\n        [\n            {\n                \"split\": \"Calibration\",\n                \"policy\": \"calibrated_thresholds\",\n                \"selected\": selected_decision_policy == \"calibrated_thresholds\",\n                **calibration_candidate_metrics,\n            },\n            {\n                \"split\": \"Validation\",\n                \"policy\": \"calibrated_thresholds\",\n                \"selected\": selected_decision_policy == \"calibrated_thresholds\",\n                **val_candidate_metrics,\n            },\n        ]\n    )\ndecision_policy_comparison = pd.DataFrame(decision_policy_rows)\nbest_epoch_frame = pd.DataFrame([checkpoint.get(\"history_row\", {})])\nvalidation_metrics_frame = pd.DataFrame([val_metrics])\nhardest_validation_cases = val_prediction_frame.sort_values([\"distance\", \"confidence\"], ascending=[False, True]).head(12).copy()\nhardest_validation_cases[\"true_class\"] = hardest_validation_cases[\"true_label\"].map(\n    lambda class_idx: CLASS_NAMES[int(class_idx)]\n)\nhardest_validation_cases[\"pred_class\"] = hardest_validation_cases[\"pred_label\"].map(\n    lambda class_idx: CLASS_NAMES[int(class_idx)]\n)\n\ndisplay(calibration_report.round(4))\ndisplay(decision_threshold_frame.round(4))\ndisplay(decision_policy_comparison.round(4))\nprint(\"Validation TTA:\", CFG.tta_horizontal)\nprint(\"Learned temperature:\", f\"{temperature:.4f}\")\nprint(\"Selected decision policy:\", selected_decision_policy)\nif decision_thresholds is not None:\n    print(\"Applied decision thresholds:\", [f\"{threshold:.4f}\" for threshold in decision_thresholds])\nelse:\n    print(\"Applied decision thresholds:\", \"argmax retained\")\ndisplay(best_epoch_frame.round(4))\ndisplay(validation_metrics_frame.round(4))\ndisplay(class_report_frame.round(4))\ndisplay(hardest_validation_cases)\n\ntraining_history_fig = plot_training_history(history_df)\nconfusion_figure, confusion_count_frame, confusion_normalized_frame = plot_confusion_heatmap(\n    val_prediction_frame[\"true_label\"].to_numpy(),\n    val_prediction_frame[\"pred_label\"].to_numpy(),\n    CLASS_NAMES,\n)\nroc_figure = plot_multiclass_roc(\n    val_prediction_frame[\"true_label\"].to_numpy(),\n    val_probabilities,\n    CLASS_NAMES,\n)\nreliability_figure, reliability_bins = plot_reliability_diagram(\n    val_prediction_frame[\"true_label\"].to_numpy(),\n    val_probabilities,\n)\n\nsave_figure(training_history_fig, \"training_history.png\")\nsave_figure(confusion_figure, \"validation_confusion_matrix.png\")\nsave_figure(roc_figure, \"validation_roc_curves.png\")\nsave_figure(reliability_figure, \"validation_reliability_diagram.png\")\n\nexport_table_bundle(history_df, \"training_history\", include_index=False)\nexport_table_bundle(best_epoch_frame, \"best_epoch_metrics\", include_index=False)\nexport_table_bundle(validation_metrics_frame, \"validation_metrics\", include_index=False)\nexport_table_bundle(class_report_frame, \"validation_classification_report\", include_index=True)\nexport_table_bundle(calibration_report, \"calibration_report\", include_index=False)\nexport_table_bundle(decision_threshold_frame, \"decision_thresholds\", include_index=False)\nexport_table_bundle(decision_policy_comparison, \"decision_policy_comparison\", include_index=False)\nexport_table_bundle(confusion_count_frame, \"validation_confusion_matrix_counts\", include_index=True)\nexport_table_bundle(confusion_normalized_frame, \"validation_confusion_matrix_normalized\", include_index=True)\nexport_table_bundle(reliability_bins, \"validation_reliability_bins\", include_index=False)\nexport_table_bundle(hardest_validation_cases, \"validation_hard_cases\", include_index=False)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:06:18.195052Z","iopub.execute_input":"2026-04-09T09:06:18.195538Z","iopub.status.idle":"2026-04-09T09:07:24.910560Z","shell.execute_reply.started":"2026-04-09T09:06:18.195502Z","shell.execute_reply":"2026-04-09T09:07:24.909660Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"59f35c7b","cell_type":"markdown","source":"## Confidence-Based Referral\n\nThis section implements the proposal idea of **referring uncertain borderline cases** instead of forcing a risky class decision.\n\nTo keep the analysis more research-accurate, the referral thresholds are **selected on the calibration split** and only then applied to the validation split.\n\nA sample is referred when either:\n\n- the top probability is too low, or\n- the margin between the top two classes is too small **and** those two classes are adjacent grades","metadata":{}},{"id":"c5c2368b","cell_type":"code","source":"def build_referral_mask(\n    probabilities: np.ndarray,\n    confidence_threshold: float,\n    margin_threshold: float,\n) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:\n    ranking = np.argsort(probabilities, axis=1)[:, ::-1]\n    top1 = ranking[:, 0]\n    top2 = ranking[:, 1]\n    top1_conf = probabilities[np.arange(len(probabilities)), top1]\n    top2_conf = probabilities[np.arange(len(probabilities)), top2]\n    margins = top1_conf - top2_conf\n    adjacent_borderline = np.abs(top1 - top2) == 1\n    referred = (top1_conf < confidence_threshold) | ((margins < margin_threshold) & adjacent_borderline)\n    return referred, top1_conf, margins, top2\n\n\ndef evaluate_referral_policy(\n    y_true: np.ndarray,\n    probabilities: np.ndarray,\n    confidence_threshold: float,\n    margin_threshold: float,\n) -> pd.Series:\n    referred, confidences, margins, runner_up = build_referral_mask(\n        probabilities,\n        confidence_threshold=confidence_threshold,\n        margin_threshold=margin_threshold,\n    )\n    predictions = probabilities.argmax(axis=1)\n    kept = ~referred\n\n    kept_metrics = compute_metrics(y_true[kept], predictions[kept], probabilities[kept], CFG.num_classes)\n    return pd.Series(\n        {\n            \"coverage\": float(kept.mean()),\n            \"referral_rate\": float(referred.mean()),\n            \"kept_accuracy\": kept_metrics[\"accuracy\"],\n            \"kept_macro_f1\": kept_metrics[\"macro_f1\"],\n            \"kept_qwk\": kept_metrics[\"qwk\"],\n            \"avg_confidence\": float(confidences.mean()),\n            \"avg_margin\": float(margins.mean()),\n        }\n    )\n\n\ndef sweep_referral_grid(y_true: np.ndarray, probabilities: np.ndarray) -> pd.DataFrame:\n    rows = []\n    for confidence_threshold in np.linspace(0.45, 0.80, 8):\n        for margin_threshold in np.linspace(0.05, 0.20, 4):\n            row = evaluate_referral_policy(y_true, probabilities, confidence_threshold, margin_threshold).to_dict()\n            row[\"confidence_threshold\"] = float(confidence_threshold)\n            row[\"margin_threshold\"] = float(margin_threshold)\n            rows.append(row)\n    return pd.DataFrame(rows)\n\n\ndef select_referral_operating_point(referral_grid: pd.DataFrame, min_coverage: float) -> pd.Series:\n    candidate_rows = referral_grid[referral_grid[\"coverage\"] >= min_coverage].copy()\n    if candidate_rows.empty:\n        candidate_rows = referral_grid.copy()\n    return candidate_rows.sort_values(\n        [\"kept_qwk\", \"kept_macro_f1\", \"coverage\"],\n        ascending=False,\n    ).iloc[0]\n\n\nreferral_tuning_frame = calibration_prediction_frame if len(calibration_prediction_frame) else val_prediction_frame\nreferral_tuning_probabilities = calibration_probabilities if len(calibration_prediction_frame) else val_probabilities\nreferral_selection_split = \"calibration\" if len(calibration_prediction_frame) else \"validation_fallback\"\n\nreferral_grid = sweep_referral_grid(\n    y_true=referral_tuning_frame[\"true_label\"].to_numpy(),\n    probabilities=referral_tuning_probabilities,\n)\nbest_referral = select_referral_operating_point(referral_grid, CFG.referral_min_coverage)\n\nreferral_summary = evaluate_referral_policy(\n    y_true=val_prediction_frame[\"true_label\"].to_numpy(),\n    probabilities=val_probabilities,\n    confidence_threshold=float(best_referral[\"confidence_threshold\"]),\n    margin_threshold=float(best_referral[\"margin_threshold\"]),\n)\ndefault_referral_summary = evaluate_referral_policy(\n    y_true=val_prediction_frame[\"true_label\"].to_numpy(),\n    probabilities=val_probabilities,\n    confidence_threshold=CFG.referral_confidence_threshold,\n    margin_threshold=CFG.referral_margin_threshold,\n)\n\nreferral_comparison = pd.DataFrame(\n    [\n        {\n            \"policy\": f\"selected_on_{referral_selection_split}_applied_to_validation\",\n            **referral_summary.to_dict(),\n            \"confidence_threshold\": float(best_referral[\"confidence_threshold\"]),\n            \"margin_threshold\": float(best_referral[\"margin_threshold\"]),\n        },\n        {\n            \"policy\": \"default_thresholds_on_validation\",\n            **default_referral_summary.to_dict(),\n            \"confidence_threshold\": CFG.referral_confidence_threshold,\n            \"margin_threshold\": CFG.referral_margin_threshold,\n        },\n    ]\n)\nselected_referral_frame = pd.DataFrame([best_referral])\ndisplay(selected_referral_frame.round(4))\ndisplay(referral_comparison.round(4))\n\nreferral_tradeoff_figure = plt.figure(figsize=(8, 5))\nsns.lineplot(data=referral_grid, x=\"coverage\", y=\"kept_qwk\", hue=\"margin_threshold\", palette=\"viridis\")\nplt.scatter(\n    [float(best_referral[\"coverage\"])],\n    [float(best_referral[\"kept_qwk\"])],\n    color=\"black\",\n    s=80,\n    label=\"Selected point\",\n    zorder=5,\n)\nplt.title(\"Referral Coverage vs Retained-Case QWK\")\nplt.xlabel(\"Coverage\")\nplt.ylabel(\"QWK on Non-Referred Cases\")\nplt.legend()\nplt.tight_layout()\nsave_figure(referral_tradeoff_figure, \"referral_tradeoff.png\")\nplt.show()\n\nexport_table_bundle(referral_grid, \"referral_grid\", include_index=False)\nexport_table_bundle(selected_referral_frame, \"selected_referral_operating_point\", include_index=False)\nexport_table_bundle(referral_comparison, \"referral_comparison\", include_index=False)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:07:24.914075Z","iopub.execute_input":"2026-04-09T09:07:24.914545Z","iopub.status.idle":"2026-04-09T09:07:25.983573Z","shell.execute_reply.started":"2026-04-09T09:07:24.914500Z","shell.execute_reply":"2026-04-09T09:07:25.982922Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"45fb6be1","cell_type":"markdown","source":"## Explainability and Quantitative XAI\n\nTwo explanation mechanisms are used:\n\n- **Grad-CAM** over the EfficientNet feature maps\n- **Attention Rollout** over the DeiT transformer layers\n\nQuantitative checks implemented here:\n\n- **faithfulness**: does confidence drop when the most salient regions are removed?\n- **stability**: do explanations remain similar after a small image perturbation?\n- **localization hook**: if lesion masks become available later, IoU and Dice can be computed too","metadata":{}},{"id":"2904a0fc","cell_type":"code","source":"def denormalize_image(image_tensor: torch.Tensor) -> np.ndarray:\n    mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1)\n    std = torch.tensor(IMAGENET_STD).view(3, 1, 1)\n    image = image_tensor.detach().cpu() * std + mean\n    image = image.clamp(0, 1).permute(1, 2, 0).numpy()\n    return image\n\n\ndef overlay_heatmap(image: np.ndarray, heatmap: np.ndarray, alpha: float = 0.35) -> np.ndarray:\n    colorized = plt.cm.jet(heatmap)[..., :3]\n    return np.clip((1.0 - alpha) * image + alpha * colorized, 0, 1)\n\n\nclass AttentionRollout:\n    def __init__(self, vit_model: nn.Module):\n        self.vit_model = vit_model\n        self.attention_maps = []\n        self.handles = [block.attn.register_forward_hook(self._hook) for block in self.vit_model.blocks]\n\n    def _hook(self, module, inputs, outputs) -> None:\n        tokens = inputs[0].detach()\n        batch_size, token_count, channels = tokens.shape\n        qkv = (\n            module.qkv(tokens)\n            .reshape(batch_size, token_count, 3, module.num_heads, channels // module.num_heads)\n            .permute(2, 0, 3, 1, 4)\n        )\n        query, key, _ = qkv.unbind(0)\n        query = module.q_norm(query)\n        key = module.k_norm(key)\n        attention = (query * module.scale) @ key.transpose(-2, -1)\n        attention = attention.softmax(dim=-1)\n        self.attention_maps.append(attention.detach())\n\n    def generate(self, image_tensor: torch.Tensor) -> np.ndarray:\n        self.attention_maps.clear()\n        with torch.no_grad():\n            _ = self.vit_model.forward_features(image_tensor)\n\n        rollout = None\n        for attention in self.attention_maps:\n            avg_attention = attention.mean(dim=1)\n            identity = torch.eye(avg_attention.size(-1), device=avg_attention.device).unsqueeze(0)\n            avg_attention = avg_attention + identity\n            avg_attention = avg_attention / avg_attention.sum(dim=-1, keepdim=True)\n            rollout = avg_attention if rollout is None else avg_attention @ rollout\n\n        mask = rollout[:, 0, 1:]\n        side = int(mask.size(-1) ** 0.5)\n        mask = mask.reshape(mask.size(0), 1, side, side)\n        mask = F.interpolate(mask, size=(CFG.image_size, CFG.image_size), mode=\"bilinear\", align_corners=False)\n        mask = mask.squeeze(1)\n        mask = (mask - mask.amin(dim=(1, 2), keepdim=True)) / (mask.amax(dim=(1, 2), keepdim=True) - mask.amin(dim=(1, 2), keepdim=True) + 1e-8)\n        return mask.detach().cpu().numpy()\n\n    def close(self) -> None:\n        for handle in self.handles:\n            handle.remove()\n\n\ndef compute_gradcam(model: nn.Module, image_tensor: torch.Tensor, class_idx: int | None = None) -> tuple[np.ndarray, int, np.ndarray]:\n    model.eval()\n    if image_tensor.ndim == 3:\n        image_tensor = image_tensor.unsqueeze(0)\n    image_tensor = image_tensor.to(DEVICE)\n\n    with torch.enable_grad():\n        outputs = model(image_tensor)\n        logits = outputs[\"logits\"]\n        target_class = int(logits.argmax(dim=1).item()) if class_idx is None else int(class_idx)\n        model.zero_grad(set_to_none=True)\n        cnn_maps = outputs[\"cnn_maps\"]\n        cnn_maps.retain_grad()\n        logits[:, target_class].sum().backward()\n\n        gradients = cnn_maps.grad\n        weights = gradients.mean(dim=(2, 3), keepdim=True)\n        cam = (weights * cnn_maps).sum(dim=1)\n        cam = F.relu(cam)\n        cam = F.interpolate(cam.unsqueeze(1), size=image_tensor.shape[-2:], mode=\"bilinear\", align_corners=False)\n        cam = cam.squeeze(1)\n        cam = (cam - cam.amin(dim=(1, 2), keepdim=True)) / (cam.amax(dim=(1, 2), keepdim=True) - cam.amin(dim=(1, 2), keepdim=True) + 1e-8)\n\n    probabilities = logits.softmax(dim=1).detach().cpu().numpy()[0]\n    return cam.detach().cpu().numpy()[0], target_class, probabilities\n\n\ndef explanation_from_method(\n    model: nn.Module,\n    rollout_helper: AttentionRollout,\n    image_tensor: torch.Tensor,\n    method_name: str,\n    class_idx: int | None = None,\n) -> tuple[np.ndarray, int, np.ndarray]:\n    if method_name == \"Grad-CAM\":\n        return compute_gradcam(model, image_tensor, class_idx=class_idx)\n\n    with torch.no_grad():\n        if image_tensor.ndim == 3:\n            image_tensor = image_tensor.unsqueeze(0)\n        image_tensor = image_tensor.to(DEVICE)\n        outputs = model(image_tensor)\n        probabilities = outputs[\"logits\"].softmax(dim=1).detach().cpu().numpy()[0]\n        target_class = int(probabilities.argmax()) if class_idx is None else int(class_idx)\n        saliency = rollout_helper.generate(image_tensor)[0]\n        return saliency, target_class, probabilities\n\n\ndef delete_salient_region(image_tensor: torch.Tensor, saliency_map: np.ndarray, fraction: float = 0.15) -> torch.Tensor:\n    threshold = np.quantile(saliency_map.reshape(-1), 1.0 - fraction)\n    mask = torch.from_numpy((saliency_map >= threshold).astype(np.float32)).to(image_tensor.device)\n    blurred = F.avg_pool2d(image_tensor.unsqueeze(0), kernel_size=21, stride=1, padding=10).squeeze(0)\n    return image_tensor * (1.0 - mask.unsqueeze(0)) + blurred * mask.unsqueeze(0)\n\n\ndef localization_overlap(saliency_map: np.ndarray, lesion_mask_path: str | None = None) -> dict[str, float]:\n    if lesion_mask_path is None or not Path(lesion_mask_path).exists():\n        return {\"iou\": float(\"nan\"), \"dice\": float(\"nan\")}\n\n    lesion_mask = Image.open(lesion_mask_path).convert(\"L\").resize((saliency_map.shape[1], saliency_map.shape[0]))\n    lesion_mask = (np.array(lesion_mask) > 127).astype(np.uint8)\n    saliency_binary = (saliency_map >= np.quantile(saliency_map.reshape(-1), 0.85)).astype(np.uint8)\n\n    intersection = (saliency_binary & lesion_mask).sum()\n    union = (saliency_binary | lesion_mask).sum()\n    iou = intersection / max(union, 1)\n    dice = (2 * intersection) / max(saliency_binary.sum() + lesion_mask.sum(), 1)\n    return {\"iou\": float(iou), \"dice\": float(dice)}\n\n\ndef evaluate_saliency_suite(\n    model: nn.Module,\n    dataset: Dataset,\n    indices: list[int],\n    rollout_helper: AttentionRollout,\n    methods: tuple[str, ...] = (\"Grad-CAM\", \"Attention Rollout\"),\n) -> pd.DataFrame:\n    rows = []\n\n    for index in indices:\n        sample = dataset[index]\n        image_tensor = sample[\"image\"].unsqueeze(0).to(DEVICE)\n\n        for method_name in methods:\n            saliency_map, target_class, base_probabilities = explanation_from_method(\n                model=model,\n                rollout_helper=rollout_helper,\n                image_tensor=image_tensor,\n                method_name=method_name,\n            )\n\n            deleted = delete_salient_region(image_tensor[0], saliency_map, fraction=CFG.deletion_fraction).unsqueeze(0)\n            with torch.no_grad():\n                deleted_probabilities = model(deleted)[\"logits\"].softmax(dim=1).detach().cpu().numpy()[0]\n\n            noisy_input = image_tensor + torch.randn_like(image_tensor) * CFG.stability_noise_std\n            perturbed_map, _, _ = explanation_from_method(\n                model=model,\n                rollout_helper=rollout_helper,\n                image_tensor=noisy_input,\n                method_name=method_name,\n                class_idx=target_class,\n            )\n\n            stability = np.corrcoef(saliency_map.reshape(-1), perturbed_map.reshape(-1))[0, 1]\n            overlap = localization_overlap(saliency_map, lesion_mask_path=None)\n\n            rows.append(\n                {\n                    \"image_id\": sample[\"image_id\"],\n                    \"method\": method_name,\n                    \"target_class\": target_class,\n                    \"faithfulness_drop\": float(base_probabilities[target_class] - deleted_probabilities[target_class]),\n                    \"stability_corr\": float(stability),\n                    \"iou\": overlap[\"iou\"],\n                    \"dice\": overlap[\"dice\"],\n                }\n            )\n\n    return pd.DataFrame(rows)\n\n\ndef pick_balanced_indices(frame: pd.DataFrame, max_per_class: int = 2) -> list[int]:\n    selected = []\n    for label in range(CFG.num_classes):\n        label_indices = frame.index[frame[\"label\"] == label].tolist()[:max_per_class]\n        selected.extend(label_indices)\n    return selected\n\n\ndef plot_explanation_gallery(\n    model: nn.Module,\n    dataset: Dataset,\n    indices: list[int],\n    rollout_helper: AttentionRollout,\n):\n    if not indices:\n        print(\"No samples available for explanation plotting.\")\n        return None\n\n    fig, axes = plt.subplots(len(indices), 3, figsize=(12, 4 * len(indices)))\n    if len(indices) == 1:\n        axes = np.expand_dims(axes, axis=0)\n\n    for row_idx, index in enumerate(indices):\n        sample = dataset[index]\n        image_tensor = sample[\"image\"]\n        gradcam_map, gradcam_target, gradcam_probs = compute_gradcam(model, image_tensor.unsqueeze(0), class_idx=None)\n        rollout_map, rollout_target, rollout_probs = explanation_from_method(\n            model=model,\n            rollout_helper=rollout_helper,\n            image_tensor=image_tensor.unsqueeze(0),\n            method_name=\"Attention Rollout\",\n            class_idx=gradcam_target,\n        )\n\n        image_rgb = denormalize_image(image_tensor)\n        axes[row_idx, 0].imshow(image_rgb)\n        axes[row_idx, 0].set_title(\n            f\"Image\\ntrue={CLASS_NAMES[int(sample['label'])]} | pred={CLASS_NAMES[int(gradcam_probs.argmax())]}\"\n        )\n        axes[row_idx, 1].imshow(overlay_heatmap(image_rgb, gradcam_map))\n        axes[row_idx, 1].set_title(\"Grad-CAM\")\n        axes[row_idx, 2].imshow(overlay_heatmap(image_rgb, rollout_map))\n        axes[row_idx, 2].set_title(\"Attention Rollout\")\n\n        for axis in axes[row_idx]:\n            axis.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n    return fig","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:07:25.984806Z","iopub.execute_input":"2026-04-09T09:07:25.985468Z","iopub.status.idle":"2026-04-09T09:07:26.017753Z","shell.execute_reply.started":"2026-04-09T09:07:25.985436Z","shell.execute_reply":"2026-04-09T09:07:26.017050Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d99ca813","cell_type":"code","source":"rollout_helper = AttentionRollout(model.vit)\n\ngallery_indices = pick_balanced_indices(val_df, max_per_class=1)[: CFG.gradcam_samples]\nexplanation_gallery_figure = plot_explanation_gallery(model, val_dataset, gallery_indices, rollout_helper)\nif explanation_gallery_figure is not None:\n    save_figure(explanation_gallery_figure, \"xai_explanation_gallery.png\")\n\nxai_indices = pick_balanced_indices(val_df, max_per_class=max(1, CFG.xai_eval_samples // CFG.num_classes))\nxai_results = evaluate_saliency_suite(model, val_dataset, xai_indices, rollout_helper)\nxai_summary_frame = (\n    xai_results.groupby(\"method\")[[\"faithfulness_drop\", \"stability_corr\", \"iou\", \"dice\"]]\n    .mean()\n    .reset_index()\n)\ndisplay(xai_summary_frame.round(4))\ndisplay(xai_results.head())\nexport_table_bundle(xai_summary_frame, \"xai_quantitative_summary\", include_index=False)\nexport_table_bundle(xai_results, \"xai_sample_results\", include_index=False)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:07:26.018935Z","iopub.execute_input":"2026-04-09T09:07:26.019273Z","iopub.status.idle":"2026-04-09T09:07:51.906580Z","shell.execute_reply.started":"2026-04-09T09:07:26.019247Z","shell.execute_reply":"2026-04-09T09:07:51.905890Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"3c156d30","cell_type":"markdown","source":"## Test Inference and Submission Export\n\nThis final section runs the best saved model on the APTOS test set, exports a Kaggle-style submission file, and packages the paper-ready figures, tables, and notebook outputs into a single downloadable zip bundle.","metadata":{}},{"id":"9426c76f","cell_type":"code","source":"def is_cuda_oom(error: RuntimeError) -> bool:\n    error_text = str(error)\n    return \"CUDA out of memory\" in error_text or \"CUDA error: out of memory\" in error_text\n\n\ndef predict_test_logits_with_fallback(\n    model: nn.Module,\n    images: torch.Tensor,\n    tta_horizontal: bool = False,\n) -> torch.Tensor:\n    try:\n        with get_autocast_context():\n            return predict_logits_with_tta(model, images, tta_horizontal=tta_horizontal)\n    except RuntimeError as error:\n        if not torch.cuda.is_available() or not is_cuda_oom(error) or images.size(0) <= 1:\n            raise\n        torch.cuda.empty_cache()\n        midpoint = max(1, images.size(0) // 2)\n        left_logits = predict_test_logits_with_fallback(model, images[:midpoint], tta_horizontal=tta_horizontal)\n        right_logits = predict_test_logits_with_fallback(model, images[midpoint:], tta_horizontal=tta_horizontal)\n        return torch.cat([left_logits, right_logits], dim=0)\n\n\ndef collect_unlabeled_logits(\n    model: nn.Module,\n    loader: DataLoader,\n    tta_horizontal: bool = False,\n) -> tuple[pd.DataFrame, np.ndarray]:\n    model.eval()\n    rows = []\n    all_logits = []\n\n    with torch.inference_mode():\n        for batch in loader:\n            images = prepare_batch_images(batch[\"image\"])\n            logits_tensor = predict_test_logits_with_fallback(model, images, tta_horizontal=tta_horizontal)\n            logits = logits_tensor.cpu().numpy()\n\n            for image_id in batch[\"image_id\"]:\n                rows.append(\n                    {\n                        \"image_id\": image_id,\n                    }\n                )\n            all_logits.append(logits)\n\n            del images, logits_tensor\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n\n    if not all_logits:\n        return pd.DataFrame(columns=[\"image_id\"]), np.empty((0, CFG.num_classes), dtype=np.float32)\n    return pd.DataFrame(rows), np.concatenate(all_logits)\n\n\ndef build_unlabeled_prediction_frame(\n    metadata_frame: pd.DataFrame,\n    logits: np.ndarray,\n    temperature: float = 1.0,\n    decision_thresholds: list[float] | np.ndarray | None = None,\n) -> tuple[pd.DataFrame, np.ndarray]:\n    probabilities = logits_to_probabilities(logits, temperature=temperature)\n    expected_grade = expected_grade_from_probabilities(probabilities)\n    if decision_thresholds is None:\n        predicted_labels = probabilities.argmax(axis=1).astype(int)\n        decision_policy = \"argmax\"\n    else:\n        predicted_labels = predict_from_expected_grade(expected_grade, decision_thresholds)\n        decision_policy = \"calibrated_thresholds\"\n\n    predicted_confidences = probabilities[np.arange(len(predicted_labels)), predicted_labels].astype(np.float32)\n    prediction_frame = metadata_frame.copy()\n    prediction_frame[\"diagnosis\"] = predicted_labels\n    prediction_frame[\"confidence\"] = predicted_confidences.astype(float)\n    prediction_frame[\"expected_grade\"] = expected_grade.astype(float)\n    prediction_frame[\"decision_policy\"] = decision_policy\n    return prediction_frame, probabilities\n\n\nif \"rollout_helper\" in globals():\n    rollout_helper.attention_maps.clear()\n    rollout_helper.close()\n    del rollout_helper\nmodel.zero_grad(set_to_none=True)\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n\naptos_test_dataset = RetinopathyDataset(\n    aptos_test_frame,\n    eval_transform,\n    CFG,\n    cache_dir=preprocessed_cache_root / \"test\" if preprocessed_cache_root is not None else None,\n)\ntest_batch_size = min(CFG.batch_size, 8 if CFG.tta_horizontal else 16)\naptos_test_loader = DataLoader(\n    aptos_test_dataset,\n    batch_size=test_batch_size,\n    shuffle=False,\n    **loader_kwargs,\n)\nprint(\"Test inference batch size:\", test_batch_size)\n\ntest_metadata_frame, test_logits = collect_unlabeled_logits(\n    model,\n    aptos_test_loader,\n    tta_horizontal=CFG.tta_horizontal,\n)\nsubmission_frame, test_probabilities = build_unlabeled_prediction_frame(\n    test_metadata_frame,\n    test_logits,\n    temperature=temperature,\n    decision_thresholds=decision_thresholds,\n)\nsubmission_path = WORKDIR / \"submission_hybrid_dr.csv\"\nsubmission_export_frame = submission_frame[[\"image_id\", \"diagnosis\"]].rename(columns={\"image_id\": \"id_code\"})\nsubmission_export_frame.to_csv(submission_path, index=False)\nregister_artifact(submission_path)\n\nexport_table_bundle(val_prediction_frame, \"validation_predictions\", include_index=False)\nexport_table_bundle(submission_export_frame, \"submission_hybrid_dr\", include_index=False)\nexport_json_artifact(asdict(CFG), \"run_config.json\")\n\npaper_key_result_rows = [\n    {\n        \"analysis\": \"validation_overall\",\n        \"accuracy\": val_metrics[\"accuracy\"],\n        \"macro_f1\": val_metrics[\"macro_f1\"],\n        \"qwk\": val_metrics[\"qwk\"],\n        \"macro_auc_ovr\": val_metrics[\"macro_auc_ovr\"],\n        \"ece\": val_metrics[\"ece\"],\n        \"coverage\": 1.0,\n        \"referral_rate\": 0.0,\n        \"faithfulness_drop\": float(\"nan\"),\n        \"stability_corr\": float(\"nan\"),\n    }\n]\nfor _, referral_row in referral_comparison.iterrows():\n    paper_key_result_rows.append(\n        {\n            \"analysis\": str(referral_row[\"policy\"]),\n            \"accuracy\": float(referral_row[\"kept_accuracy\"]),\n            \"macro_f1\": float(referral_row[\"kept_macro_f1\"]),\n            \"qwk\": float(referral_row[\"kept_qwk\"]),\n            \"macro_auc_ovr\": float(\"nan\"),\n            \"ece\": float(\"nan\"),\n            \"coverage\": float(referral_row[\"coverage\"]),\n            \"referral_rate\": float(referral_row[\"referral_rate\"]),\n            \"faithfulness_drop\": float(\"nan\"),\n            \"stability_corr\": float(\"nan\"),\n        }\n    )\nfor _, xai_row in xai_summary_frame.iterrows():\n    paper_key_result_rows.append(\n        {\n            \"analysis\": f\"xai_{xai_row['method']}\",\n            \"accuracy\": float(\"nan\"),\n            \"macro_f1\": float(\"nan\"),\n            \"qwk\": float(\"nan\"),\n            \"macro_auc_ovr\": float(\"nan\"),\n            \"ece\": float(\"nan\"),\n            \"coverage\": float(\"nan\"),\n            \"referral_rate\": float(\"nan\"),\n            \"faithfulness_drop\": float(xai_row[\"faithfulness_drop\"]),\n            \"stability_corr\": float(xai_row[\"stability_corr\"]),\n        }\n    )\npaper_key_results_frame = pd.DataFrame(paper_key_result_rows)\nexport_table_bundle(paper_key_results_frame, \"paper_key_results\", include_index=False)\n\nartifact_descriptions = {\n    \"figures/dataset_distribution.png\": \"Dataset and source distribution figure for the data section.\",\n    \"figures/training_history.png\": \"Training and validation learning curves.\",\n    \"figures/validation_confusion_matrix.png\": \"Validation confusion matrix figure with counts and row-normalized views.\",\n    \"figures/validation_roc_curves.png\": \"One-vs-rest ROC curves for the validation split.\",\n    \"figures/validation_reliability_diagram.png\": \"Reliability diagram after temperature scaling.\",\n    \"figures/referral_tradeoff.png\": \"Coverage versus retained-case QWK referral figure.\",\n    \"figures/xai_explanation_gallery.png\": \"Qualitative Grad-CAM and Attention Rollout gallery.\",\n    \"tables/dataset_overall_class_counts.csv\": \"Overall dataset class-count table.\",\n    \"tables/dataset_source_class_counts.csv\": \"Per-source class-count table.\",\n    \"tables/split_class_distribution.csv\": \"Train, calibration, and validation split counts by class.\",\n    \"tables/split_source_label_distribution.csv\": \"Per-source split counts by class.\",\n    \"tables/training_history.csv\": \"Per-epoch optimization history.\",\n    \"tables/best_epoch_metrics.csv\": \"Metrics from the checkpoint-selected best epoch.\",\n    \"tables/validation_metrics.csv\": \"Main validation metrics for the paper results table.\",\n    \"tables/validation_predictions.csv\": \"Per-sample validation predictions and confidences.\",\n    \"tables/validation_classification_report.csv\": \"Precision, recall, and F1 by class.\",\n    \"tables/decision_thresholds.csv\": \"Ordinal thresholds tuned on the calibration split.\",\n    \"tables/decision_policy_comparison.csv\": \"Argmax versus calibrated-threshold performance comparison.\",\n    \"tables/validation_confusion_matrix_counts.csv\": \"Raw validation confusion matrix counts.\",\n    \"tables/validation_confusion_matrix_normalized.csv\": \"Row-normalized validation confusion matrix.\",\n    \"tables/validation_reliability_bins.csv\": \"Reliability-diagram bin statistics.\",\n    \"tables/validation_hard_cases.csv\": \"Most difficult validation examples for error analysis.\",\n    \"tables/calibration_report.csv\": \"Calibration before and after temperature scaling.\",\n    \"tables/referral_grid.csv\": \"Referral-threshold sweep grid.\",\n    \"tables/selected_referral_operating_point.csv\": \"Chosen referral operating point.\",\n    \"tables/referral_comparison.csv\": \"Comparison of selected and default referral policies.\",\n    \"tables/xai_quantitative_summary.csv\": \"Mean faithfulness and stability by XAI method.\",\n    \"tables/xai_sample_results.csv\": \"Per-image XAI evaluation results.\",\n    \"tables/submission_hybrid_dr.csv\": \"Kaggle-format submission table.\",\n    \"tables/paper_key_results.csv\": \"Compact paper summary table across validation, referral, and XAI.\",\n    \"reports/run_config.json\": \"Notebook configuration used for the run.\",\n    \"submission_hybrid_dr.csv\": \"Submission CSV written to the working directory root.\",\n}\nmanifest_lines = [\n    \"# Paper Output Manifest\",\n    \"\",\n    \"This bundle contains manuscript-ready figures, tables, and reproducibility files generated by the notebook.\",\n    \"\",\n    \"## Included Artifacts\",\n]\nfor artifact_path in sorted(EXPORTED_ARTIFACTS, key=lambda path: str(path)):\n    try:\n        relative_path = artifact_path.relative_to(ARTIFACT_DIR)\n        relative_text = str(relative_path).replace(\"\\\\\", \"/\")\n    except ValueError:\n        relative_text = artifact_path.name\n    description = artifact_descriptions.get(relative_text, \"Notebook-generated output artifact.\")\n    manifest_lines.append(f\"- {relative_text}: {description}\")\nexport_text_artifact(\"\\n\".join(manifest_lines) + \"\\n\", \"paper_output_manifest.md\")\n\ndownload_bundle_path = ARTIFACT_DIR / \"hybrid_dr_notebook_outputs.zip\"\nwith zipfile.ZipFile(download_bundle_path, \"w\", compression=zipfile.ZIP_DEFLATED) as archive:\n    for artifact_path in sorted(EXPORTED_ARTIFACTS, key=lambda path: str(path)):\n        if artifact_path == download_bundle_path:\n            continue\n        try:\n            arcname = artifact_path.relative_to(ARTIFACT_DIR)\n        except ValueError:\n            arcname = Path(artifact_path.name)\n        archive.write(artifact_path, arcname=str(arcname).replace(\"\\\\\", \"/\"))\n\ndisplay(submission_frame.head())\ndisplay(FileLink(str(download_bundle_path), result_html_prefix=\"Download artifact bundle: \"))\ndisplay(FileLink(str(submission_path), result_html_prefix=\"Download submission CSV: \"))\nprint(\"Validation TTA used for export:\", CFG.tta_horizontal)\nprint(\"Applied temperature scaling to confidence output:\", f\"{temperature:.4f}\")\nprint(\"Applied decision policy for export:\", selected_decision_policy)\nif decision_thresholds is not None:\n    print(\"Applied calibrated ordinal thresholds:\", [f\"{threshold:.4f}\" for threshold in decision_thresholds])\nelse:\n    print(\"Applied calibrated ordinal thresholds:\", \"not used (argmax retained)\")\nprint(\"Submission written to:\", submission_path)\nprint(\"Artifact bundle written to:\", download_bundle_path)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:07:51.907735Z","iopub.execute_input":"2026-04-09T09:07:51.908037Z","iopub.status.idle":"2026-04-09T09:09:32.869595Z","shell.execute_reply.started":"2026-04-09T09:07:51.908009Z","shell.execute_reply":"2026-04-09T09:09:32.868637Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"733bf9a0","cell_type":"markdown","source":"## What To Try Next\n\n- run multi-seed experiments and report mean plus standard deviation for QWK and macro F1\n- keep a true external test set if you plan to claim cross-dataset generalization\n- compare against single-branch CNN and transformer baselines to make the hybrid contribution easier to justify\n- add lesion masks later if you want IoU/Dice-based localization scoring\n- extend the calibration study with classwise ECE or confidence histograms if the report needs deeper uncertainty analysis","metadata":{}},{"id":"8acbeeb9","cell_type":"code","source":"cache_dir = globals().get(\"preprocessed_cache_root\", WORKDIR / \"preprocessed_cache\")\ncache_dir = Path(cache_dir)\n\nif cache_dir.exists():\n    shutil.rmtree(cache_dir)\n    print(f\"Removed preprocessed cache: {cache_dir}\")\nelse:\n    print(f\"No preprocessed cache found at: {cache_dir}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T09:09:32.871562Z","iopub.execute_input":"2026-04-09T09:09:32.872013Z","iopub.status.idle":"2026-04-09T09:09:33.144728Z","shell.execute_reply.started":"2026-04-09T09:09:32.871973Z","shell.execute_reply":"2026-04-09T09:09:33.144025Z"},"trusted":true},"outputs":[],"execution_count":null}]}