{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"be94c20a","cell_type":"markdown","source":"# VinBigData Chest X-ray — Multi-Label ViT Classification\n\nEnd-to-end pipeline: DICOM loading → windowing → normalization → ViT (HuggingFace `transformers`) multi-label training.\n\n**Pipeline summary**\n- Labels collapsed from box-level `train.csv` to image-level multi-hot vectors (14 finding classes, `No finding` implied by all-zero row)\n- DICOM → windowing (per-image `WindowCenter`/`WindowWidth` from metadata, with a percentile-based fallback when metadata is missing/inconsistent) → resize to 224×224 → stack to 3 channels → normalize using **mean/std computed from a sample of our own windowed training images** (not ImageNet stats — windowed X-ray pixels don't follow a natural-photo distribution)\n- 90/10 **multi-label stratified** split (iterative stratification, so rare classes are represented proportionally in both splits)\n- `ViTForImageClassification` (`google/vit-base-patch16-224-in21k`) fine-tuned with `BCEWithLogitsLoss`\n- **Every epoch**: validation loss (checkpoint criterion) → per-class AUC → per-class best F1 threshold (grid search) → per-class F1 at that threshold → macro/micro summaries\n- Best model saved by **lowest validation loss**\n\n**Before running on Kaggle**\n- Add the VinBigData Chest X-ray Abnormalities Detection dataset (DICOM version) under Add Data\n- Enable a GPU accelerator (Settings → Accelerator → GPU T4 x2 or P100)\n- Confirm dataset paths in the Config cell match your attached dataset's directory name (Kaggle mounts datasets under `/kaggle/input/<dataset-slug>/`, and the slug can vary — printed for you in the path-check cell)\n","metadata":{}},{"id":"fbf8ab92","cell_type":"markdown","source":"## 1. Setup","metadata":{}},{"id":"1a556b6e","cell_type":"code","source":"# Kaggle ships most of these, but pin/install what's missing or version-sensitive.\n# pydicom: DICOM reading. iterative-stratification: multi-label stratified split.\nimport sys, subprocess\n\ndef pip_install(pkg):\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pkg])\n\nfor pkg in [\"pydicom\", \"iterative-stratification\", \"transformers\", \"accelerate\"]:\n    pip_install(pkg)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:47.538018Z","iopub.execute_input":"2026-06-30T06:29:47.538768Z","iopub.status.idle":"2026-06-30T06:29:59.921823Z","shell.execute_reply.started":"2026-06-30T06:29:47.538721Z","shell.execute_reply":"2026-06-30T06:29:59.921122Z"}},"outputs":[],"execution_count":null},{"id":"017e1b3e","cell_type":"code","source":"import os\nimport json\nimport math\nimport random\nimport logging\nimport warnings\nfrom pathlib import Path\nfrom dataclasses import dataclass, field\nfrom typing import List, Tuple, Dict, Optional\n\nimport numpy as np\nimport pandas as pd\nimport cv2\n# import pydicom\n# from pydicom.pixel_data_handlers.util import apply_voi_lut\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import roc_auc_score, f1_score\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\n\nfrom transformers import ViTForImageClassification, ViTImageProcessor\n\nimport matplotlib.pyplot as plt\n\nwarnings.filterwarnings(\"ignore\")\n\nlogging.basicConfig(\n    level=logging.INFO,\n    format=\"%(asctime)s | %(levelname)s | %(message)s\",\n    datefmt=\"%H:%M:%S\",\n)\nlogger = logging.getLogger(\"vinbig_vit\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:59.923219Z","iopub.execute_input":"2026-06-30T06:29:59.923638Z","iopub.status.idle":"2026-06-30T06:29:59.930596Z","shell.execute_reply.started":"2026-06-30T06:29:59.923615Z","shell.execute_reply":"2026-06-30T06:29:59.92971Z"}},"outputs":[],"execution_count":null},{"id":"698b5c81","cell_type":"code","source":"def set_seed(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nlogger.info(f\"Using device: {device}\")\nif device.type == \"cuda\":\n    logger.info(f\"GPU: {torch.cuda.get_device_name(0)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:59.931629Z","iopub.execute_input":"2026-06-30T06:29:59.931907Z","iopub.status.idle":"2026-06-30T06:29:59.948243Z","shell.execute_reply.started":"2026-06-30T06:29:59.93187Z","shell.execute_reply":"2026-06-30T06:29:59.947617Z"}},"outputs":[],"execution_count":null},{"id":"8e4c7002","cell_type":"markdown","source":"## 2. Config\n\nAdjust `DATASET_DIR` if your attached Kaggle dataset has a different slug. The cell below lists `/kaggle/input` so you can confirm the exact path before training starts — silent path mismatches are the #1 cause of \"works fine, just trains on nothing.\"\n","metadata":{}},{"id":"d78c7e99","cell_type":"code","source":"print(\"Contents of /kaggle/input:\")\nfor p in sorted(Path(\"/kaggle/input\").glob(\"*\")):\n    print(\" -\", p)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:59.950281Z","iopub.execute_input":"2026-06-30T06:29:59.950506Z","iopub.status.idle":"2026-06-30T06:29:59.960642Z","shell.execute_reply.started":"2026-06-30T06:29:59.950488Z","shell.execute_reply":"2026-06-30T06:29:59.959924Z"}},"outputs":[],"execution_count":null},{"id":"b126d1b9","cell_type":"code","source":"@dataclass\nclass Config:\n    # ---- Paths (EDIT THESE if your dataset slug differs from what's printed above) ----\n    DATASET_DIR: str = \"/kaggle/input/datasets/xhlulu/vinbigdata-chest-xray-resized-png-1024x1024\"\n    TRAIN_CSV: str = field(init=False)\n    # TRAIN_DICOM_DIR: str = field(init=False)\n    \n    # ---- Output ----\n    OUTPUT_DIR: str = \"/kaggle/working\"\n    BEST_MODEL_PATH: str = field(init=False)\n    METRICS_LOG_PATH: str = field(init=False)\n\n    # ---- Classes ----\n    # 14 official finding classes. \"No finding\" (class_id 14) is NOT a column —\n    # it's implicitly represented as an all-zero label vector.\n    CLASS_NAMES: tuple = (\n        \"Aortic enlargement\", \"Atelectasis\", \"Calcification\", \"Cardiomegaly\",\n        \"Consolidation\", \"ILD\", \"Infiltration\", \"Lung Opacity\",\n        \"Nodule/Mass\", \"Other lesion\", \"Pleural effusion\", \"Pleural thickening\",\n        \"Pneumothorax\", \"Pulmonary fibrosis\",\n    )\n    NUM_CLASSES: int = 14\n    NO_FINDING_CLASS_ID: int = 14\n\n    # ---- Image preprocessing ----\n    IMG_SIZE: int = 224\n    # NORM_MEAN/NORM_STD are placeholders here — they get OVERWRITTEN with\n    # values computed from a sample of the actual windowed training images\n    # (see \"Compute dataset-specific normalization stats\" section below).\n    # We don't use ImageNet's (0.485,0.456,0.406)/(0.229,0.224,0.225) because\n    # those describe natural RGB photos; windowed chest X-rays have a very\n    # different intensity distribution, so stats computed on our own\n    # windowed pixels are a better match for what the network actually sees.\n    NORM_MEAN: tuple = (0.5, 0.5, 0.5)\n    NORM_STD: tuple = (0.5, 0.5, 0.5)\n    # How many TRAINING images (never validation — would leak val info into\n    # the normalization applied to train) to sample when estimating mean/std.\n    # A few hundred is enough for a stable estimate without re-reading the\n    # whole dataset; raise this if you want more precision at the cost of time.\n    NORM_STATS_SAMPLE_SIZE: int = 400\n    NORM_STATS_SEED: int = 42\n    # Fallback window width/center (in rescaled HU-like units) used only when\n    # DICOM metadata has no usable WindowCenter/WindowWidth.\n    FALLBACK_WINDOW_PERCENTILES: tuple = (0.5, 99.5)\n\n    # ---- Split ----\n    VAL_SIZE: float = 0.10\n    SPLIT_SEED: int = 42\n\n    # ---- Model ----\n    MODEL_NAME: str = \"google/vit-base-patch16-224-in21k\"\n\n    # ---- Training ----\n    BATCH_SIZE: int = 16\n    NUM_WORKERS: int = 6\n    EPOCHS: int = 15\n    LR: float = 3e-5\n    WEIGHT_DECAY: float = 1e-4\n    WARMUP_RATIO: float = 0.05\n    GRAD_CLIP_NORM: float = 1.0\n\n    # ---- Threshold search ----\n    THRESHOLD_GRID: tuple = tuple(np.round(np.arange(0.05, 0.96, 0.01), 2))\n\n    def __post_init__(self):\n        self.TRAIN_CSV = f\"/kaggle/input/competitions/vinbigdata-chest-xray-abnormalities-detection/train.csv\"\n        self.TRAIN_PNG_DIR = f\"/kaggle/input/datasets/xhlulu/vinbigdata-chest-xray-resized-png-1024x1024/train\"\n        self.BEST_MODEL_PATH = f\"{self.OUTPUT_DIR}/best_vit_vinbig.pt\"\n        self.METRICS_LOG_PATH = f\"{self.OUTPUT_DIR}/metrics_per_epoch.json\"\n\n\ncfg = Config()\nos.makedirs(cfg.OUTPUT_DIR, exist_ok=True)\nlogger.info(f\"Train CSV: {cfg.TRAIN_CSV}\")\nlogger.info(f\"PNG dir: {cfg.TRAIN_PNG_DIR}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:59.961811Z","iopub.execute_input":"2026-06-30T06:29:59.962237Z","iopub.status.idle":"2026-06-30T06:29:59.980241Z","shell.execute_reply.started":"2026-06-30T06:29:59.962217Z","shell.execute_reply":"2026-06-30T06:29:59.979689Z"}},"outputs":[],"execution_count":null},{"id":"68493b9f","cell_type":"markdown","source":"## 3. Build image-level multi-label table\n\n`train.csv` is **box-level**: each row is one annotated finding (or one \"No finding\" row with no box). The same `image_id` appears multiple times when several findings are present. We collapse this to one row per `image_id` with a 14-dim multi-hot vector — `1` if that class appears *anywhere* among that image's boxes, else `0`. An image with only a \"No finding\" row, or no rows in the 14 classes, ends up all-zero, which is the correct multi-label encoding (findings absent), not a 15th column.\n","metadata":{}},{"id":"45b80773","cell_type":"code","source":"raw_df = pd.read_csv(cfg.TRAIN_CSV)\nlogger.info(f\"Raw train.csv shape (box-level): {raw_df.shape}\")\nraw_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:29:59.981131Z","iopub.execute_input":"2026-06-30T06:29:59.981404Z","iopub.status.idle":"2026-06-30T06:30:00.096828Z","shell.execute_reply.started":"2026-06-30T06:29:59.981386Z","shell.execute_reply":"2026-06-30T06:30:00.096032Z"}},"outputs":[],"execution_count":null},{"id":"4f32b155","cell_type":"code","source":"def build_image_level_labels(raw: pd.DataFrame, cfg: Config) -> pd.DataFrame:\n    \"\"\"Collapse box-level annotations into one multi-hot row per image_id.\"\"\"\n    image_ids = raw[\"image_id\"].unique()\n    label_matrix = pd.DataFrame(\n        0, index=image_ids, columns=list(cfg.CLASS_NAMES), dtype=np.int64\n    )\n\n    # Only rows with a real finding class (excludes \"No finding\" id 14, which\n    # correctly leaves that image's row as all-zero).\n    finding_rows = raw[raw[\"class_id\"] < cfg.NUM_CLASSES]\n    for image_id, class_id in zip(finding_rows[\"image_id\"], finding_rows[\"class_id\"]):\n        label_matrix.loc[image_id, cfg.CLASS_NAMES[class_id]] = 1\n\n    label_df = label_matrix.reset_index().rename(columns={\"index\": \"image_id\"})\n    return label_df\n\n\nlabel_df = build_image_level_labels(raw_df, cfg)\nlogger.info(f\"Image-level label table shape: {label_df.shape}\")\n\nclass_counts = label_df[list(cfg.CLASS_NAMES)].sum().sort_values(ascending=False)\nno_finding_count = (label_df[list(cfg.CLASS_NAMES)].sum(axis=1) == 0).sum()\nlogger.info(f\"'No finding' images (all-zero rows): {no_finding_count}\")\nprint(class_counts)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:30:00.097904Z","iopub.execute_input":"2026-06-30T06:30:00.098209Z","iopub.status.idle":"2026-06-30T06:30:01.611863Z","shell.execute_reply.started":"2026-06-30T06:30:00.098177Z","shell.execute_reply":"2026-06-30T06:30:01.611211Z"}},"outputs":[],"execution_count":null},{"id":"fe021420","cell_type":"code","source":"fig, ax = plt.subplots(figsize=(10, 5))\nclass_counts.plot(kind=\"bar\", ax=ax)\nax.set_title(\"Positive image count per class (image-level, multi-label)\")\nax.set_ylabel(\"Number of images\")\nplt.xticks(rotation=60, ha=\"right\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-30T06:30:01.612742Z","iopub.execute_input":"2026-06-30T06:30:01.613068Z"}},"outputs":[],"execution_count":null},{"id":"94ae69bd","cell_type":"markdown","source":"## 4. Multi-label stratified 90/10 split\n\nPlain `train_test_split(stratify=y)` doesn't accept a 2D multi-label `y`. We use `MultilabelStratifiedShuffleSplit` from `iterative-stratification`, which balances the *per-class* positive rate across splits rather than treating each unique label combination as one stratum — important here since rare classes (e.g. `Pneumothorax`) could otherwise be nearly absent from one split.\n","metadata":{}},{"id":"22f96f0e","cell_type":"code","source":"X_dummy = np.zeros((len(label_df), 1))  # splitter only needs y for stratification\ny_all = label_df[list(cfg.CLASS_NAMES)].values\n\nmss = MultilabelStratifiedShuffleSplit(\n    n_splits=1, test_size=cfg.VAL_SIZE, random_state=cfg.SPLIT_SEED\n)\ntrain_idx, val_idx = next(mss.split(X_dummy, y_all))\n\ntrain_df = label_df.iloc[train_idx].reset_index(drop=True)\nval_df = label_df.iloc[val_idx].reset_index(drop=True)\n\nlogger.info(f\"Train: {len(train_df)} images | Val: {len(val_df)} images \"\n            f\"({len(val_df) / len(label_df):.1%})\")\n\n# Sanity check: per-class positive RATE should be close between splits.\nrate_compare = pd.DataFrame({\n    \"train_rate\": train_df[list(cfg.CLASS_NAMES)].mean(),\n    \"val_rate\": val_df[list(cfg.CLASS_NAMES)].mean(),\n})\nrate_compare[\"abs_diff\"] = (rate_compare[\"train_rate\"] - rate_compare[\"val_rate\"]).abs()\nprint(rate_compare.sort_values(\"abs_diff\", ascending=False))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"e33e84fc","cell_type":"markdown","source":"## 5. DICOM loading and windowing\n\n**Windowing** maps the (often 12–16 bit) raw DICOM pixel range down to the diagnostically relevant intensity band a radiologist would actually look at — e.g. for chest X-rays this controls how soft tissue vs. bone vs. air is rendered. We try, in order:\n\n1. `apply_voi_lut` using the DICOM's own `WindowCenter`/`WindowWidth` (or VOI LUT sequence) — this is what the acquiring scanner/software intended.\n2. If those tags are missing, multi-valued in a way that can't be resolved, or produce a degenerate (near-constant) image, fall back to a **percentile-based window**: clip to the `[0.5, 99.5]` percentile of pixel intensities, which is a robust general-purpose substitute that avoids the all-black/all-white failure mode of trusting absent or bad metadata.\n\nAfter windowing, pixels are scaled to `[0, 1]`, resized to `IMG_SIZE`, and stacked to 3 channels (the backbone expects RGB; grayscale chest X-rays are duplicated across channels, which is standard practice for adapting pretrained color-image backbones to medical grayscale imaging). Normalization to zero-mean/unit-std happens next, in section 6, using stats computed from our own data rather than ImageNet's.\n","metadata":{}},{"id":"e35480f1","cell_type":"code","source":"# def read_dicom_pixels(path: str) -> Tuple[np.ndarray, pydicom.Dataset]:\n#     \"\"\"Read raw pixel array and dataset (for windowing metadata).\"\"\"\n#     ds = pydicom.dcmread(path)\n#     arr = ds.pixel_array.astype(np.float32)\n\n#     # Some VinBigData DICOMs are MONOCHROME1 (inverted: high value = dark).\n#     # Flip so higher value always means brighter/denser tissue, matching\n#     # MONOCHROME2 convention that windowing math below assumes.\n#     if getattr(ds, \"PhotometricInterpretation\", \"MONOCHROME2\") == \"MONOCHROME1\":\n#         arr = arr.max() - arr\n\n#     # Apply rescale slope/intercept if present (raw stored value -> real units).\n#     slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n#     intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n#     arr = arr * slope + intercept\n\n#     return arr, ds\n\n\n# def get_window_center_width(ds: pydicom.Dataset) -> Optional[Tuple[float, float]]:\n#     \"\"\"Extract a single (center, width) pair from DICOM metadata, if usable.\"\"\"\n#     wc = getattr(ds, \"WindowCenter\", None)\n#     ww = getattr(ds, \"WindowWidth\", None)\n#     if wc is None or ww is None:\n#         return None\n#     # These tags can be multi-valued (one pair per VOI LUT option) — take the first.\n#     if isinstance(wc, pydicom.multival.MultiValue):\n#         wc = float(wc[0])\n#     else:\n#         wc = float(wc)\n#     if isinstance(ww, pydicom.multival.MultiValue):\n#         ww = float(ww[0])\n#     else:\n#         ww = float(ww)\n#     if ww <= 0:\n#         return None\n#     return wc, ww\n\n\n# def apply_windowing(arr: np.ndarray, ds: pydicom.Dataset, cfg: Config) -> np.ndarray:\n#     \"\"\"Window the pixel array to [0, 1] using DICOM metadata, with a robust fallback.\"\"\"\n#     cw = get_window_center_width(ds)\n\n#     if cw is not None:\n#         center, width = cw\n#         lo, hi = center - width / 2.0, center + width / 2.0\n#         windowed = np.clip(arr, lo, hi)\n#         denom = hi - lo\n#         # Degenerate metadata (near-zero range, or windowed image collapses to a\n#         # near-constant value) -> treat as unusable and fall through to percentile method.\n#         if denom > 1e-6:\n#             windowed01 = (windowed - lo) / denom\n#             if windowed01.std() > 1e-3:\n#                 return windowed01.astype(np.float32)\n\n#     # Fallback: percentile-based windowing.\n#     lo_pct, hi_pct = cfg.FALLBACK_WINDOW_PERCENTILES\n#     lo, hi = np.percentile(arr, [lo_pct, hi_pct])\n#     if hi - lo < 1e-6:\n#         hi = arr.max()\n#         lo = arr.min() if arr.max() > arr.min() else arr.min() - 1.0\n#     windowed = np.clip(arr, lo, hi)\n#     windowed01 = (windowed - lo) / (hi - lo + 1e-8)\n#     return windowed01.astype(np.float32)\n\n\n# def window_and_resize_dicom(path: str, cfg: Config) -> np.ndarray:\n#     \"\"\"Read -> window -> resize -> stack to 3 channels. Output stays in [0, 1].\n\n#     Split out from normalization so the SAME function can be reused both for\n#     final preprocessing (below) and for sampling pixels to compute\n#     dataset-specific normalization stats, without duplicating the windowing logic.\n\n#     Returns a (3, IMG_SIZE, IMG_SIZE) float32 array, values in [0, 1].\n#     \"\"\"\n#     arr, ds = read_dicom_pixels(path)\n#     windowed01 = apply_windowing(arr, ds, cfg)  # [0, 1], single channel\n\n#     resized = cv2.resize(\n#         windowed01, (cfg.IMG_SIZE, cfg.IMG_SIZE), interpolation=cv2.INTER_LINEAR\n#     )\n#     rgb = np.stack([resized, resized, resized], axis=0)  # (3, H, W), still [0, 1]\n#     return rgb.astype(np.float32)\n\n\n# def load_and_preprocess_dicom(path: str, cfg: Config) -> np.ndarray:\n#     \"\"\"Full pipeline: read -> window -> resize -> 3-channel -> normalize.\n\n#     Normalization uses cfg.NORM_MEAN/NORM_STD, which by this point in the\n#     notebook have been OVERWRITTEN with dataset-computed values (not the\n#     ImageNet placeholder defaults) — see the normalization-stats section.\n\n#     Returns a (3, IMG_SIZE, IMG_SIZE) float32 array ready for the model.\n#     \"\"\"\n#     rgb01 = window_and_resize_dicom(path, cfg)\n\n#     mean = np.array(cfg.NORM_MEAN, dtype=np.float32).reshape(3, 1, 1)\n#     std = np.array(cfg.NORM_STD, dtype=np.float32).reshape(3, 1, 1)\n#     normalized = (rgb01 - mean) / std\n\n#     return normalized.astype(np.float32)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"d6f2210b","cell_type":"markdown","source":"## 6. Compute dataset-specific normalization stats\n\nRather than normalizing with ImageNet's mean/std, we compute mean and std **directly from a sample of our own windowed training images** — windowed chest X-ray pixels don't follow the same intensity distribution as natural RGB photos, so stats estimated from the actual data the model will see are a better fit.\n\nTwo things matter for correctness here:\n- **Train-only**: we sample exclusively from `train_df`, never `val_df`. Computing stats on validation images (even just to estimate normalization) would let validation data influence how the model preprocesses training data — a subtle form of leakage.\n- **Post-windowing**: we sample from the *windowed, resized* `[0, 1]` pixel output (via `window_and_resize_dicom`), not raw DICOM values — raw pixel intensities are on an arbitrary device-dependent scale and windowing is exactly the transform that puts them somewhere meaningful, so stats must be computed after it, not before.\n\nSince the windowed image is stacked identically across all 3 channels (grayscale duplicated for the RGB-expecting backbone), the 3 channels have identical statistics by construction — so we compute one scalar mean/std from the grayscale values and broadcast it across all 3 channels, rather than computing 3 redundant (and noisier, since each would be estimated independently) per-channel values.\n","metadata":{}},{"id":"9bcc9dd4","cell_type":"code","source":"# # Visual sanity check: windowing should produce images with visible anatomical\n# # contrast, not flat gray/black/white. Spot-check a few before committing to training.\n# # Uses window_and_resize_dicom directly (skips normalization) so this check is\n# # purely about windowing quality, independent of normalization stats.\n# sample_ids = train_df[\"image_id\"].sample(4, random_state=0).tolist()\n\n# fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n# for ax, image_id in zip(axes, sample_ids):\n#     dicom_path = f\"{cfg.TRAIN_DICOM_DIR}/{image_id}.dicom\"\n#     windowed_rgb01 = window_and_resize_dicom(dicom_path, cfg)  # [0, 1], 3 identical channels\n#     ax.imshow(windowed_rgb01[0], cmap=\"gray\", vmin=0, vmax=1)\n#     ax.set_title(image_id[:12] + \"...\", fontsize=9)\n#     ax.axis(\"off\")\n# plt.suptitle(\"Post-windowing samples (pre-normalization, [0,1] range)\")\n# plt.tight_layout()\n# plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"28729fa4-c8a5-4c88-b4d5-5c83469a72e3","cell_type":"markdown","source":"## 5. PNG loading, Resizing and Normalization","metadata":{}},{"id":"9349f40b-0a75-433f-8018-358e7207c8b7","cell_type":"code","source":"def read_png_pixels(path):\n    arr = cv2.imread(path, cv2.IMREAD_GRAYSCALE).astype(np.float32)\n    return arr\n\ndef load_and_resize_png(path, cfg):\n    arr = read_png_pixels(path)          # [0,255] غالبًا\n    normalized01 = arr / 255.0           # أو apply_windowing لو محتاج fallback\n    resized = cv2.resize(normalized01, (cfg.IMG_SIZE, cfg.IMG_SIZE), interpolation=cv2.INTER_LINEAR)\n    rgb = np.stack([resized, resized, resized], axis=0)\n    return rgb.astype(np.float32)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"9739c144-8d2f-4d06-828a-756357c2b6d7","cell_type":"code","source":"# Visual sanity check: windowing should produce images with visible anatomical\n# contrast, not flat gray/black/white. Spot-check a few before committing to training.\n# Uses window_and_resize_dicom directly (skips normalization) so this check is\n# purely about windowing quality, independent of normalization stats.\nsample_ids = train_df[\"image_id\"].sample(4, random_state=0).tolist()\n\nfig, axes = plt.subplots(1, 4, figsize=(16, 4))\nfor ax, image_id in zip(axes, sample_ids):\n    png_path = f\"{cfg.TRAIN_PNG_DIR}/{image_id}.png\"\n    windowed_rgb01 = load_and_resize_png(png_path, cfg)  # [0, 1], 3 identical channels\n    ax.imshow(windowed_rgb01[0], cmap=\"gray\", vmin=0, vmax=1)\n    ax.set_title(image_id[:12] + \"...\", fontsize=9)\n    ax.axis(\"off\")\nplt.suptitle(\"Post-windowing samples (pre-normalization, [0,1] range)\")\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"4ba8a3e3","cell_type":"code","source":"def compute_normalization_stats(df: pd.DataFrame, dicom_dir: str, cfg: Config) -> Tuple[float, float]:\n    \"\"\"Sample N training images, run them through windowing+resize, and compute\n    a single scalar (mean, std) over all sampled pixels.\n\n    Returns (mean, std) as plain floats, to be broadcast across all 3 channels.\n    \"\"\"\n    rng = np.random.RandomState(cfg.NORM_STATS_SEED)\n    sample_size = min(cfg.NORM_STATS_SAMPLE_SIZE, len(df))\n    sampled_ids = rng.choice(df[\"image_id\"].values, size=sample_size, replace=False)\n\n    pixel_sum = 0.0\n    pixel_sq_sum = 0.0\n    pixel_count = 0\n\n    for image_id in sampled_ids:\n        dicom_path = f\"{dicom_dir}/{image_id}.png\"\n        windowed_rgb01 = load_and_resize_png(dicom_path, cfg)  # (3, H, W), [0, 1]\n        single_channel = windowed_rgb01[0]  # channels are identical, use one\n\n        pixel_sum += single_channel.sum(dtype=np.float64)\n        pixel_sq_sum += np.square(single_channel, dtype=np.float64).sum()\n        pixel_count += single_channel.size\n\n    mean = pixel_sum / pixel_count\n    variance = (pixel_sq_sum / pixel_count) - (mean ** 2)\n    std = math.sqrt(max(variance, 1e-12))  # guard against tiny negative from float error\n\n    return float(mean), float(std)\n\n\nlogger.info(f\"Computing normalization stats from {cfg.NORM_STATS_SAMPLE_SIZE} sampled TRAINING images \"\n            f\"(train_df only, post-windowing)...\")\ncomputed_mean, computed_std = compute_normalization_stats(train_df, cfg.TRAIN_PNG_DIR, cfg)\n\nlogger.info(f\"Computed mean: {computed_mean:.4f} | Computed std: {computed_std:.4f}\")\nlogger.info(f\"(For comparison, ImageNet single-channel-equivalent is roughly mean=0.449, std=0.226)\")\n\n# Overwrite cfg with the dataset-computed stats, broadcast across all 3 channels.\n# Everything downstream (window_and_resize_dicom callers, the Dataset class\n# defined next, load_and_preprocess_dicom) reads cfg.NORM_MEAN/NORM_STD, so this\n# single assignment is enough to apply the new stats everywhere.\ncfg.NORM_MEAN = (computed_mean, computed_mean, computed_mean)\ncfg.NORM_STD = (computed_std, computed_std, computed_std)\n\nlogger.info(f\"cfg.NORM_MEAN set to {cfg.NORM_MEAN}\")\nlogger.info(f\"cfg.NORM_STD set to {cfg.NORM_STD}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"8705b150","cell_type":"code","source":"# Sanity check: with the new stats, normalized pixels from a few sample images\n# should land roughly in [-2, 2] (a handful of values further out is expected for\n# image extremes), with mean near 0 and std near 1 — confirming the stats actually\n# describe the data they were computed from before we commit to using them.\ncheck_ids = train_df[\"image_id\"].sample(20, random_state=1).tolist()\nall_normalized_pixels = []\n\nfor image_id in check_ids:\n    dicom_path = f\"{cfg.TRAIN_PNG_DIR}/{image_id}.png\"\n    normalized = load_and_resize_png(dicom_path, cfg)\n    all_normalized_pixels.append(normalized[0].ravel())\n\nall_normalized_pixels = np.concatenate(all_normalized_pixels)\nlogger.info(f\"Post-normalization check ({len(check_ids)} held-out training images): \"\n            f\"mean={all_normalized_pixels.mean():.4f}, std={all_normalized_pixels.std():.4f}, \"\n            f\"min={all_normalized_pixels.min():.4f}, max={all_normalized_pixels.max():.4f}\")\n\nfig, ax = plt.subplots(figsize=(8, 4))\nax.hist(all_normalized_pixels, bins=100)\nax.axvline(0, color=\"red\", linestyle=\"--\", alpha=0.6, label=\"0 (target mean)\")\nax.set_title(\"Distribution of normalized pixel values (sample check)\")\nax.set_xlabel(\"Normalized pixel value\")\nax.legend()\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"fc8ae5ac","cell_type":"markdown","source":"## 7. Dataset and DataLoader","metadata":{}},{"id":"6470f88c","cell_type":"code","source":"class VinBigDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, dicom_dir: str, cfg: Config, augment: bool = False):\n        self.df = df.reset_index(drop=True)\n        self.dicom_dir = dicom_dir\n        self.cfg = cfg\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id = row[\"image_id\"]\n        dicom_path = f\"{self.dicom_dir}/{image_id}.png\"\n\n        image = load_and_resize_png(dicom_path, self.cfg)\n\n        if self.augment:\n            image = self._augment(image)\n\n        labels = row[list(self.cfg.CLASS_NAMES)].values.astype(np.float32)\n\n        return {\n            \"pixel_values\": torch.from_numpy(image).float(),\n            \"labels\": torch.from_numpy(labels).float(),\n            \"image_id\": image_id,\n        }\n\n    def _augment(self, image: np.ndarray) -> np.ndarray:\n        # Light, anatomy-safe augmentation only — chest X-rays are orientation-\n        # sensitive (no vertical flips: that would mirror left/right anatomy\n        # inconsistently with the label, e.g. cardiomegaly laterality cues).\n        if random.random() < 0.5:\n            image = np.ascontiguousarray(image[:, :, ::-1])  # horizontal flip\n        if random.random() < 0.3:\n            noise = np.random.normal(0, 0.01, image.shape).astype(np.float32)\n            image = image + noise\n        return image\n\n\ntrain_dataset = VinBigDataset(train_df, cfg.TRAIN_PNG_DIR, cfg, augment=True)\nval_dataset = VinBigDataset(val_df, cfg.TRAIN_PNG_DIR, cfg, augment=False)\n\ntrain_loader = DataLoader(\n    train_dataset, batch_size=cfg.BATCH_SIZE, shuffle=True,\n    num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True,\n)\nval_loader = DataLoader(\n    val_dataset, batch_size=cfg.BATCH_SIZE, shuffle=False,\n    num_workers=cfg.NUM_WORKERS, pin_memory=True,\n)\n\nlogger.info(f\"Train batches: {len(train_loader)} | Val batches: {len(val_loader)}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"a3cffdd9","cell_type":"markdown","source":"## 8. Model\n\n`ViTForImageClassification` with `problem_type=\"multi_label_classification\"`, which makes the HF model internally use `BCEWithLogitsLoss` when labels are passed as float multi-hot vectors — we also define the loss explicitly below so the training loop is self-contained and not dependent on that internal default.\n","metadata":{}},{"id":"83e2f2d1","cell_type":"code","source":"model = ViTForImageClassification.from_pretrained(\n    cfg.MODEL_NAME,\n    num_labels=cfg.NUM_CLASSES,\n    problem_type=\"multi_label_classification\",\n    ignore_mismatched_sizes=True,  # replaces the pretrained head with a fresh NUM_CLASSES head\n)\nmodel.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\n\nlogger.info(f\"Model: {cfg.MODEL_NAME}\")\nlogger.info(f\"Total params: {sum(p.numel() for p in model.parameters()):,}\")\nlogger.info(f\"Trainable params: {sum(p.numel() for p in model.parameters() if p.requires_grad):,}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"d9be3587","cell_type":"code","source":"optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n\ntotal_steps = len(train_loader) * cfg.EPOCHS\nwarmup_steps = int(total_steps * cfg.WARMUP_RATIO)\n\ndef lr_lambda(current_step):\n    if current_step < warmup_steps:\n        return current_step / max(1, warmup_steps)\n    progress = (current_step - warmup_steps) / max(1, total_steps - warmup_steps)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))  # cosine decay after warmup\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\nlogger.info(f\"Total steps: {total_steps} | Warmup steps: {warmup_steps}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"7fa20356","cell_type":"markdown","source":"## 9. Per-class AUC, best-F1 threshold search, and per-class F1\n\nThese run on the **full validation set's** predicted probabilities after each epoch (not per-batch — thresholds and AUC need the whole distribution to be meaningful):\n\n- **AUC per class**: `roc_auc_score(y_true_class, y_prob_class)`. Skipped (`NaN`) for a class if validation has only one label value present (a fold with zero positives or zero negatives for that class makes AUC undefined) — this can happen with the rarer classes despite stratification, especially early on with a small val set.\n- **Best threshold per class**: grid search over `THRESHOLD_GRID`, picking the threshold that maximizes that class's F1 on validation. This is independent per class since a global single threshold would systematically under- or over-predict rare vs. common classes.\n- **F1 per class**: F1 at that class's own best threshold.\n","metadata":{}},{"id":"a9e3489d","cell_type":"code","source":"def compute_per_class_auc(y_true: np.ndarray, y_prob: np.ndarray, class_names) -> Dict[str, float]:\n    \"\"\"y_true, y_prob: (N, C). Returns {class_name: auc}, NaN where undefined.\"\"\"\n    aucs = {}\n    for i, name in enumerate(class_names):\n        col_true = y_true[:, i]\n        if len(np.unique(col_true)) < 2:\n            aucs[name] = float(\"nan\")\n        else:\n            aucs[name] = roc_auc_score(col_true, y_prob[:, i])\n    return aucs\n\n\ndef find_best_threshold_per_class(\n    y_true: np.ndarray, y_prob: np.ndarray, class_names, grid\n) -> Tuple[Dict[str, float], Dict[str, float]]:\n    \"\"\"Grid search threshold maximizing F1, independently per class.\n\n    Returns (best_threshold_per_class, best_f1_per_class).\n    \"\"\"\n    grid = np.asarray(grid)\n    best_thresholds, best_f1s = {}, {}\n\n    for i, name in enumerate(class_names):\n        col_true = y_true[:, i]\n        col_prob = y_prob[:, i]\n\n        if col_true.sum() == 0:\n            # No positives in val for this class -> F1 is undefined/zero regardless\n            # of threshold; record default threshold 0.5 and F1 0.0 rather than NaN,\n            # since \"no positive predictions wanted\" is a meaningful, well-defined state.\n            best_thresholds[name] = 0.5\n            best_f1s[name] = 0.0\n            continue\n\n        f1_scores = np.array([\n            f1_score(col_true, (col_prob >= t).astype(int), zero_division=0)\n            for t in grid\n        ])\n        best_idx = int(np.argmax(f1_scores))\n        best_thresholds[name] = float(grid[best_idx])\n        best_f1s[name] = float(f1_scores[best_idx])\n\n    return best_thresholds, best_f1s\n\n\ndef summarize_epoch_metrics(\n    y_true: np.ndarray, y_prob: np.ndarray, class_names, grid\n) -> Dict:\n    \"\"\"One-call wrapper producing everything needed for a single epoch's report.\"\"\"\n    auc_per_class = compute_per_class_auc(y_true, y_prob, class_names)\n    best_thresholds, f1_per_class = find_best_threshold_per_class(y_true, y_prob, class_names, grid)\n\n    auc_values = [v for v in auc_per_class.values() if not math.isnan(v)]\n    macro_auc = float(np.mean(auc_values)) if auc_values else float(\"nan\")\n    macro_f1 = float(np.mean(list(f1_per_class.values())))\n\n    y_pred_binary = np.zeros_like(y_prob)\n    for i, name in enumerate(class_names):\n        y_pred_binary[:, i] = (y_prob[:, i] >= best_thresholds[name]).astype(int)\n    micro_f1 = f1_score(y_true, y_pred_binary, average=\"micro\", zero_division=0)\n\n    return {\n        \"auc_per_class\": auc_per_class,\n        \"best_threshold_per_class\": best_thresholds,\n        \"f1_per_class\": f1_per_class,\n        \"macro_auc\": macro_auc,\n        \"macro_f1\": macro_f1,\n        \"micro_f1\": micro_f1,\n    }\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"c545357a","cell_type":"markdown","source":"## 10. Train and validation epoch functions","metadata":{}},{"id":"0f06346b","cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scheduler, criterion, device, grad_clip_norm):\n    model.train()\n    running_loss = 0.0\n\n    for batch in loader:\n        pixel_values = batch[\"pixel_values\"].to(device, non_blocking=True)\n        labels = batch[\"labels\"].to(device, non_blocking=True)\n\n        optimizer.zero_grad()\n        outputs = model(pixel_values=pixel_values)\n        logits = outputs.logits\n        loss = criterion(logits, labels)\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)\n        optimizer.step()\n        scheduler.step()\n\n        running_loss += loss.item() * pixel_values.size(0)\n\n    return running_loss / len(loader.dataset)\n\n\n@torch.no_grad()\ndef validate_one_epoch(model, loader, criterion, device, num_classes):\n    model.eval()\n    running_loss = 0.0\n    all_probs = np.zeros((len(loader.dataset), num_classes), dtype=np.float32)\n    all_labels = np.zeros((len(loader.dataset), num_classes), dtype=np.float32)\n\n    cursor = 0\n    for batch in loader:\n        pixel_values = batch[\"pixel_values\"].to(device, non_blocking=True)\n        labels = batch[\"labels\"].to(device, non_blocking=True)\n\n        outputs = model(pixel_values=pixel_values)\n        logits = outputs.logits\n        loss = criterion(logits, labels)\n        running_loss += loss.item() * pixel_values.size(0)\n\n        probs = torch.sigmoid(logits).cpu().numpy()\n        bsz = probs.shape[0]\n        all_probs[cursor:cursor + bsz] = probs\n        all_labels[cursor:cursor + bsz] = labels.cpu().numpy()\n        cursor += bsz\n\n    val_loss = running_loss / len(loader.dataset)\n    return val_loss, all_labels, all_probs\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"18f215f5","cell_type":"markdown","source":"## 11. Main training loop\n\nCheckpointing rule: save whenever validation loss is the lowest seen so far. Per-class AUC/threshold/F1 are computed every epoch regardless of whether that epoch is a new best, so you get the full trajectory in `metrics_per_epoch.json` even for non-improving epochs.\n","metadata":{}},{"id":"925e9b2b","cell_type":"code","source":"history = []\nbest_val_loss = float(\"inf\")\nbest_epoch = -1\n\nfor epoch in range(1, cfg.EPOCHS + 1):\n    logger.info(f\"--- Epoch {epoch}/{cfg.EPOCHS} ---\")\n\n    train_loss = train_one_epoch(\n        model, train_loader, optimizer, scheduler, criterion, device, cfg.GRAD_CLIP_NORM\n    )\n    val_loss, y_true, y_prob = validate_one_epoch(model, val_loader, criterion, device, cfg.NUM_CLASSES)\n\n    epoch_metrics = summarize_epoch_metrics(y_true, y_prob, cfg.CLASS_NAMES, cfg.THRESHOLD_GRID)\n\n    logger.info(f\"Train loss: {train_loss:.4f} | Val loss: {val_loss:.4f}\")\n    macro_auc_val = epoch_metrics[\"macro_auc\"]\n    macro_f1_val = epoch_metrics[\"macro_f1\"]\n    micro_f1_val = epoch_metrics[\"micro_f1\"]\n    logger.info(f\"Macro AUC: {macro_auc_val:.4f} | \"\n                f\"Macro F1: {macro_f1_val:.4f} | \"\n                f\"Micro F1: {micro_f1_val:.4f}\")\n\n    per_class_table = pd.DataFrame({\n        \"AUC\": epoch_metrics[\"auc_per_class\"],\n        \"Best_Threshold\": epoch_metrics[\"best_threshold_per_class\"],\n        \"F1\": epoch_metrics[\"f1_per_class\"],\n    })\n    print(per_class_table.round(4))\n\n    record = {\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"val_loss\": val_loss,\n        **epoch_metrics,\n    }\n    history.append(record)\n\n    # Persist metrics after every epoch so a crash/disconnect doesn't lose history.\n    with open(cfg.METRICS_LOG_PATH, \"w\") as f:\n        json.dump(history, f, indent=2)\n\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        best_epoch = epoch\n        torch.save({\n            \"epoch\": epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"val_loss\": val_loss,\n            \"best_threshold_per_class\": epoch_metrics[\"best_threshold_per_class\"],\n            \"class_names\": list(cfg.CLASS_NAMES),\n            \"config\": cfg.__dict__,\n        }, cfg.BEST_MODEL_PATH)\n        logger.info(f\"New best model saved (val_loss={val_loss:.4f}) -> {cfg.BEST_MODEL_PATH}\")\n\nlogger.info(f\"Training complete. Best epoch: {best_epoch} (val_loss={best_val_loss:.4f})\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"717a2554","cell_type":"markdown","source":"## 12. Training curves","metadata":{}},{"id":"e6dc88d8","cell_type":"code","source":"history_df = pd.DataFrame([\n    {\"epoch\": h[\"epoch\"], \"train_loss\": h[\"train_loss\"], \"val_loss\": h[\"val_loss\"],\n     \"macro_auc\": h[\"macro_auc\"], \"macro_f1\": h[\"macro_f1\"], \"micro_f1\": h[\"micro_f1\"]}\n    for h in history\n])\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(history_df[\"epoch\"], history_df[\"train_loss\"], label=\"Train loss\")\naxes[0].plot(history_df[\"epoch\"], history_df[\"val_loss\"], label=\"Val loss\")\naxes[0].axvline(best_epoch, color=\"gray\", linestyle=\"--\", alpha=0.6, label=\"Best epoch\")\naxes[0].set_title(\"Loss\")\naxes[0].set_xlabel(\"Epoch\")\naxes[0].legend()\n\naxes[1].plot(history_df[\"epoch\"], history_df[\"macro_auc\"], label=\"Macro AUC\", color=\"green\")\naxes[1].set_title(\"Macro AUC\")\naxes[1].set_xlabel(\"Epoch\")\naxes[1].legend()\n\naxes[2].plot(history_df[\"epoch\"], history_df[\"macro_f1\"], label=\"Macro F1\")\naxes[2].plot(history_df[\"epoch\"], history_df[\"micro_f1\"], label=\"Micro F1\")\naxes[2].set_title(\"F1 (at per-class best thresholds)\")\naxes[2].set_xlabel(\"Epoch\")\naxes[2].legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"30a1799e","cell_type":"code","source":"# Per-class AUC trajectory across all epochs (useful to see which classes the\n# model struggles with throughout training, not just at the final epoch).\nauc_trajectory = pd.DataFrame(\n    [h[\"auc_per_class\"] for h in history], index=[h[\"epoch\"] for h in history]\n)\n\nfig, ax = plt.subplots(figsize=(12, 6))\nfor class_name in cfg.CLASS_NAMES:\n    ax.plot(auc_trajectory.index, auc_trajectory[class_name], label=class_name, alpha=0.8)\nax.set_xlabel(\"Epoch\")\nax.set_ylabel(\"AUC\")\nax.set_title(\"Per-class AUC across training\")\nax.legend(bbox_to_anchor=(1.02, 1), loc=\"upper left\", fontsize=8)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"fe6ff254","cell_type":"markdown","source":"## 13. Best-epoch summary","metadata":{}},{"id":"cd583592","cell_type":"code","source":"best_record = next(h for h in history if h[\"epoch\"] == best_epoch)\n\nprint(f\"Best epoch: {best_epoch}\")\nprint(f\"Validation loss: {best_record['val_loss']:.4f}\")\nprint(f\"Macro AUC: {best_record['macro_auc']:.4f}\")\nprint(f\"Macro F1: {best_record['macro_f1']:.4f}  |  Micro F1: {best_record['micro_f1']:.4f}\")\nprint()\n\nfinal_table = pd.DataFrame({\n    \"AUC\": best_record[\"auc_per_class\"],\n    \"Best_Threshold\": best_record[\"best_threshold_per_class\"],\n    \"F1\": best_record[\"f1_per_class\"],\n}).round(4)\nfinal_table\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"c8fe0ec3","cell_type":"code","source":"# Confirm the saved checkpoint matches what we expect, and show how to reload it\n# for inference later (e.g. on the VinBigData test set).\ncheckpoint = torch.load(cfg.BEST_MODEL_PATH, map_location=device)\nprint(f\"Loaded checkpoint from epoch {checkpoint['epoch']}, val_loss={checkpoint['val_loss']:.4f}\")\nprint(f\"Saved at: {cfg.BEST_MODEL_PATH}\")\nprint(f\"Full metrics history saved at: {cfg.METRICS_LOG_PATH}\")\n\n# Example reload pattern for inference:\n# inference_model = ViTForImageClassification.from_pretrained(\n#     cfg.MODEL_NAME, num_labels=cfg.NUM_CLASSES, problem_type=\"multi_label_classification\"\n# )\n# inference_model.load_state_dict(checkpoint[\"model_state_dict\"])\n# inference_model.to(device).eval()\n# thresholds = checkpoint[\"best_threshold_per_class\"]  # apply per-class at inference time\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}