{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"Shared CNN training pipeline for GramSwasthya's image-based Phase 1 diseases.\n\nDiabetic retinopathy, chest X-ray and skin lesion all differ in exactly three\nthings: which folder the images live in, what the classes are called, and which\nbackbone to fine-tune. Everything else — discovering the data, splitting it,\nhandling class imbalance, scoring, and writing the four artifacts the backend\nloads — is identical, so it lives here and each disease script is a few lines of\nconfiguration.\n\nFour artifacts are written per run, and the backend needs all four:\n\n    model.pt           self-describing checkpoint: weights + backbone + classes\n    labels.json        index -> class name, in the order the model emits\n    metrics.json       how it scored during training, and on what data\n    preprocessing.json the exact transform to reproduce at inference time\n\nA fifth is written only when `final_test_fraction` is set:\n\n    final_test_metrics.json   scored exactly once, after training ends\n\nThe distinction between metrics.json and final_test_metrics.json is the point of\nthe split, not a filing detail. metrics.json is measured every epoch, so it is\nwhat you watch, tune against and stop on — and anything you steer by stops being\nan unbiased estimate of how the model behaves on data it has never met. The final\nslice is set aside before the first epoch and read once, after the weights are\nfrozen and written, so it can still answer that question honestly.\n\n`preprocessing.json` is the one that silently breaks things when it is missing.\nA model fine-tuned on 224px ImageNet-normalised crops will return confident\nnonsense if the backend feeds it anything else, and nothing about the checkpoint\nreveals the mismatch. Writing it next to the weights keeps the two together.\n\"\"\"\n\nfrom __future__ import annotations\n\nimport json\nimport platform\nimport random\nimport re\nimport shutil\nimport sys\nimport time\nfrom dataclasses import dataclass, field, asdict\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom PIL import Image\nfrom torch.utils.data import DataLoader, Subset, WeightedRandomSampler\nfrom torchvision import datasets, models, transforms\n\n# ImageNet statistics — every torchvision pretrained backbone expects these.\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\nSUPPORTED_BACKBONES = (\"efficientnet_b0\", \"mobilenet_v2\")\nIMAGE_SUFFIXES = {\".png\", \".jpg\", \".jpeg\", \".bmp\", \".tif\", \".tiff\", \".webp\"}\n\n# Anything zipped on a Mac carries two kinds of debris: a parallel `__MACOSX/`\n# tree mirroring the real folder layout, and an AppleDouble `._name` sidecar\n# beside every real file. Both matter here, and for different reasons.\n#\n# The sidecars end in .jpeg/.png like real images but hold resource-fork metadata,\n# so PIL raises UnidentifiedImageError the moment one reaches a DataLoader — which\n# is what killed a cxr/mobilenet_v2 run on\n# `__MACOSX/chest_xray/train/PNEUMONIA/._person122_virus_229.jpeg`.\n#\n# The `__MACOSX/` tree is the quieter danger. It mirrors `chest_xray/train/NORMAL`\n# and `.../PNEUMONIA` exactly, and every one of those folders is full of\n# image-suffixed sidecars — so a class-folder search scores it as a perfect match\n# and can train on a directory containing no actual pixels.\nMACOS_METADATA_DIR = \"__MACOSX\"\nAPPLEDOUBLE_PREFIX = \"._\"\n\n\ndef is_macos_metadata(path: str | Path) -> bool:\n    \"\"\"True for macOS zip debris: anything under __MACOSX/, or a ._ sidecar.\"\"\"\n    candidate = Path(path)\n    return (MACOS_METADATA_DIR in candidate.parts\n            or candidate.name.startswith(APPLEDOUBLE_PREFIX))\n\n\ndef report_macos_debris(count: int, where: str = \"\") -> None:\n    \"\"\"Say how much was dropped. Never silent, on purpose.\n\n    A run that quietly skips unreadable files looks identical to one that skipped a\n    genuinely corrupt image, and those need opposite responses — this is harmless\n    and expected, that is a data problem worth chasing. Printing the count keeps\n    the two distinguishable.\n    \"\"\"\n    if count:\n        print(f\"[data] excluded {count} macOS zip metadata file(s)\"\n              + (f\" under {where}\" if where else \"\")\n              + f\" ({MACOS_METADATA_DIR}/ or {APPLEDOUBLE_PREFIX}* sidecars) — \"\n              \"not images, safe to ignore\")\n\n\ndef working_root() -> Path:\n    \"\"\"Where scratch and artifacts go: /kaggle/working, or a local dir off-Kaggle.\n\n    The scripts are written for Kaggle, but falling back keeps them runnable on a\n    laptop for a smoke test — which is the difference between finding a bug now\n    and finding it twenty minutes into a GPU session.\n    \"\"\"\n    kaggle = Path(\"/kaggle/working\")\n    try:\n        kaggle.mkdir(parents=True, exist_ok=True)\n        probe = kaggle / \".write_probe\"\n        probe.touch()\n        probe.unlink()\n        return kaggle\n    except OSError:\n        local = Path.cwd() / \"_work\"\n        local.mkdir(parents=True, exist_ok=True)\n        return local\n\n\ndef rebase_to_working(path: str | Path) -> Path:\n    \"\"\"Map a /kaggle/working/... path onto working_root(); leave others alone.\"\"\"\n    candidate = Path(path)\n    try:\n        return working_root() / candidate.relative_to(\"/kaggle/working\")\n    except ValueError:\n        return candidate\n\n\n@dataclass\nclass TrainingConfig:\n    \"\"\"Everything a disease training script has to decide.\"\"\"\n\n    disease_id: str\n    dataset_dir: str\n    class_names: list[str]\n    backbone: str = \"efficientnet_b0\"\n    image_size: int = 224\n    epochs: int = 5\n    batch_size: int = 32\n    learning_rate: float = 3e-4\n    val_split: float = 0.15\n    seed: int = 42\n    num_workers: int = 2\n    # Kaggle sessions are time-boxed; capping images per class keeps a demo run\n    # to minutes on the big sets (DR is ~35k images). None means use everything.\n    max_images_per_class: int | None = 400\n    output_dir: str | None = None\n    model_version: str = \"demo-0.1.0\"\n    # Off by default: only fundus photography has a black frame worth removing,\n    # and a crop applied to a chest film or a dermoscopy image would be damage.\n    retina_crop: bool = False\n    retina_crop_threshold: int = 10\n    # Fraction of the supplied holdout to sequester as a true final test slice,\n    # scored once after training and never before. 0.0 disables it, which is the\n    # right default: it only makes sense when a holdout is supplied at all, and\n    # carving one out of an already-small validation set costs more than it buys.\n    final_test_fraction: float = 0.0\n\n    def resolved_output_dir(self) -> Path:\n        raw = (Path(self.output_dir) if self.output_dir\n               else Path(\"/kaggle/working\") / self.disease_id / self.backbone)\n        return rebase_to_working(raw)\n\n\n# --------------------------------------------------------------------------\n# Finding the data\n# --------------------------------------------------------------------------\n\ndef _normalise(name: str) -> str:\n    \"\"\"Fold a folder name to something comparable: 'No_DR' / '0 - No DR' -> 'nodr'.\"\"\"\n    return re.sub(r\"[^a-z0-9]\", \"\", name.lower())\n\n\ndef _has_images(directory: Path, minimum: int = 1) -> bool:\n    seen = 0\n    for entry in directory.iterdir():\n        if (entry.is_file() and entry.suffix.lower() in IMAGE_SUFFIXES\n                and not is_macos_metadata(entry)):\n            seen += 1\n            if seen >= minimum:\n                return True\n    return False\n\n\ndef _slug_matches(name: str, slug: str) -> bool:\n    \"\"\"Compare a folder name to a dataset slug, ignoring case and punctuation.\n\n    A self-uploaded dataset is named by whoever uploaded it, so \"DERM12345 Skin\n    Lesion\" and \"derm12345-skin-lesion\" have to compare equal.\n    \"\"\"\n    a, b = _normalise(name), _normalise(slug)\n    if not a or not b:\n        return False\n    return b in a or (len(a) >= 6 and a in b)\n\n\ndef resolve_dataset_root(slug: str, marker: str = \"\", extra_candidates=(),\n                         label: str = \"\", mount_root: str | Path = \"/kaggle/input\") -> Path | None:\n    \"\"\"Find where Kaggle actually mounted a dataset, whoever attached it.\n\n    The same dataset shows up at several different paths depending on which tab a\n    teammate attached it from and, for a self-uploaded set, which account holds it:\n\n        /kaggle/input/<slug>                     Datasets tab\n        /kaggle/input/competitions/<slug>        Competitions tab\n        /kaggle/input/datasets/<owner>/<slug>    full-path form\n\n    **The owner segment is never hardcoded.** DERM12345 is uploaded separately by\n    each teammate to their own account, so that segment differs from person to\n    person; naming any one username would fix the run for its uploader and quietly\n    break it for everybody else, who would get a synthetic-fallback run that reads\n    like a broken path rather than a missing dataset. Owner is always globbed.\n\n    `marker` is a file that must be present for a root to count. Directory\n    existence is not evidence on its own: /kaggle/input/competitions exists as soon\n    as any competition is attached, and an empty folder looks exactly like the\n    wrong one. With no marker, a non-empty directory reached via the slug counts.\n\n    Returns the root, or None having printed what is actually mounted — a wrong\n    path should cost one run and a one-line edit, not a puzzling synthetic run.\n\n    `mount_root` exists so this is testable off Kaggle. /kaggle cannot be created\n    without root, so the previous mount fix could only ever be reasoned about;\n    pointing this at a fixture tree lets the resolution logic be exercised for\n    real instead.\n    \"\"\"\n    name = label or slug\n    mounts = Path(mount_root)\n\n    def usable(path: Path) -> bool:\n        if not path.is_dir():\n            return False\n        if marker:\n            return (path / marker).exists() or any(path.glob(f\"*/{marker}\"))\n        return any(path.iterdir())\n\n    def announce(root: Path, how: str = \"\") -> Path:\n        found = f\"  ({marker} found)\" if marker else \"\"\n        print(f\"[data] {name} root: {root}{found}{how}\")\n        return root\n\n    # Tier 1: the exact conventions, callers' confirmed paths first.\n    candidates = [Path(c) for c in extra_candidates]\n    candidates += [mounts / slug, mounts / \"competitions\" / slug,\n                   mounts / \"datasets\" / slug]\n    # Tier 2: same slug, any owner. This is the line that makes a self-uploaded\n    # dataset work for a teammate who is not the uploader.\n    candidates += sorted(mounts.glob(f\"datasets/*/{slug}\"))\n    candidates += sorted(mounts.glob(f\"competitions/*/{slug}\"))\n\n    checked: list[tuple[Path, bool, bool]] = []\n    seen: set[Path] = set()\n    for root in candidates:\n        if root in seen:\n            continue\n        seen.add(root)\n        ok = usable(root)\n        checked.append((root, root.is_dir(), ok))\n        if ok:\n            return announce(root)\n\n    # Tier 3: any owner AND a differently-spelled upload, matched on the\n    # normalised slug so \"DERM12345 Skin Lesion\" still resolves.\n    for pattern in (\"datasets/*/*\", \"competitions/*/*\", \"*\"):\n        for match in sorted(mounts.glob(pattern)) if mounts.is_dir() else []:\n            if match in seen or not match.is_dir():\n                continue\n            if _slug_matches(match.name, slug) and usable(match):\n                return announce(match, \"  (matched on name, not an exact slug)\")\n\n    # Tier 4: nested deeper than any convention. Slug-guarded throughout, so\n    # another dataset's file of the same name can never be picked up by accident.\n    if mounts.is_dir():\n        for depth in (\"*/\", \"*/*/\", \"*/*/*/\", \"*/*/*/*/\"):\n            pattern = depth + marker if marker else depth\n            for match in sorted(mounts.glob(pattern)):\n                root = match.parent if marker else match\n                if root in seen or not root.is_dir():\n                    continue\n                if any(_slug_matches(part, slug) for part in root.parts) and usable(root):\n                    return announce(root, \"  (found by search — not at a standard mount)\")\n\n    print(f\"[warn] {name} not found at any known mount\"\n          + (f\", looking for {marker}\" if marker else \"\") + \":\")\n    for root, exists, ok in checked:\n        print(f\"[warn]   {root}  (directory exists: {exists}, usable: {ok})\")\n    if mounts.is_dir():\n        print(f\"[warn] {mounts} contains: \"\n              f\"{sorted(p.name for p in mounts.iterdir())}\")\n        for extra in (\"datasets\", \"competitions\"):\n            sub = mounts / extra\n            if sub.is_dir():\n                print(f\"[warn] {sub} contains: \"\n                      f\"{sorted(p.name for p in sub.iterdir())}\")\n    else:\n        print(f\"[warn] {mounts} does not exist — not running on Kaggle?\")\n    return None\n\n\ndef find_class_root(dataset_dir: Path, class_names: list[str], max_depth: int = 4):\n    \"\"\"Locate the directory whose subfolders are the classes.\n\n    Kaggle datasets nest unpredictably — the classes might sit at the top level,\n    under `train/`, or under `<dataset>/<dataset>/train/`. Rather than hardcode a\n    guess that breaks on the next dataset, walk down and score each candidate by\n    how many of the expected classes it contains. Returns (root, {class: folder})\n    or (None, {}) when nothing plausible turns up.\n    \"\"\"\n    if not dataset_dir.exists():\n        return None, {}\n\n    wanted = {_normalise(c): c for c in class_names}\n    best: tuple[int, Path, dict] = (0, dataset_dir, {})\n\n    queue: list[tuple[Path, int]] = [(dataset_dir, 0)]\n    while queue:\n        current, depth = queue.pop(0)\n        try:\n            # Sorted, because iterdir() order is filesystem-dependent: on an\n            # unsorted walk this function returned a different directory run to\n            # run for any dataset with train/val/test siblings.\n            subdirs = sorted((d for d in current.iterdir()\n                              if d.is_dir() and not is_macos_metadata(d)),\n                             key=lambda d: d.name)\n        except (PermissionError, OSError):\n            continue\n\n        mapping: dict[str, Path] = {}\n        for sub in subdirs:\n            key = _normalise(sub.name)\n            if key in wanted:\n                mapping[wanted[key]] = sub\n            else:\n                # 'Moderate' should still match a folder called '2_Moderate'.\n                for norm, original in wanted.items():\n                    if original not in mapping and norm and norm in key:\n                        mapping[original] = sub\n                        break\n\n        usable = {c: p for c, p in mapping.items() if _has_images(p)}\n        if len(usable) > best[0]:\n            best = (len(usable), current, usable)\n        if len(usable) == len(class_names):\n            # First complete match wins. When a dataset ships several split\n            # directories that each hold every class, \"first\" is a coin toss\n            # between them and this heuristic cannot know which one is training\n            # data — so say so, loudly, and let the caller pass class_dirs.\n            siblings = [d.name for d in subdirs\n                        if d != current and _normalise(d.name) not in wanted]\n            if siblings:\n                print(f\"[warn] find_class_root picked {current.name!r}, but \"\n                      f\"{siblings} sit alongside it and may also hold classes. \"\n                      f\"Pass class_dirs explicitly if the choice matters.\")\n            return current, usable\n\n        if depth < max_depth:\n            queue.extend((d, depth + 1) for d in subdirs)\n\n    if best[0] >= 2:  # a partial match is still worth training on\n        return best[1], best[2]\n    return None, {}\n\n\ndef build_synthetic_dataset(destination: Path, class_names: list[str],\n                            image_size: int, per_class: int = 24) -> Path:\n    \"\"\"Write a tiny random-noise dataset so a run completes without the real data.\n\n    This exists so the script is runnable end to end before the Kaggle dataset is\n    attached — the plumbing, the artifacts and the backend contract can all be\n    exercised on day one. Every artifact produced this way is stamped\n    `data_source: synthetic-fallback`, because a demo model trained on noise that\n    is mistaken for a real one is exactly the failure this project is built to\n    avoid.\n    \"\"\"\n    from PIL import Image\n\n    if destination.exists():\n        shutil.rmtree(destination)\n    rng = np.random.default_rng(0)\n    for index, name in enumerate(class_names):\n        class_dir = destination / name\n        class_dir.mkdir(parents=True, exist_ok=True)\n        for n in range(per_class):\n            # Give each class a different colour bias so the run is not pure noise\n            # and the loss curve actually moves; it is still meaningless data.\n            base = rng.integers(0, 60, size=(image_size, image_size, 3), dtype=np.uint8)\n            base[:, :, index % 3] = np.clip(base[:, :, index % 3] + 120, 0, 255)\n            Image.fromarray(base).save(class_dir / f\"{name}_{n:03d}.png\")\n    return destination\n\n\ndef build_class_view(mapping: dict[str, Path], class_names: list[str], destination: Path) -> Path:\n    \"\"\"Create a clean folder-per-class view, one symlink per class, in a fixed order.\n\n    Two problems disappear here. ImageFolder treats *every* subdirectory of its\n    root as a class, so a stray `models/` or `.ipynb_checkpoints` beside the real\n    classes becomes a phantom label (and an empty one raises outright on recent\n    torchvision). And ImageFolder orders labels alphabetically, which would scramble\n    CLASS_NAMES — 'Mild' before 'No_DR' before 'Severe' — so the index the backend\n    reads from labels.json would not mean what the training script intended.\n    Numbering the view directories pins both.\n    \"\"\"\n    if destination.exists():\n        shutil.rmtree(destination)\n    destination.mkdir(parents=True, exist_ok=True)\n    present = [c for c in class_names if c in mapping]\n    for index, name in enumerate(present):\n        link = destination / f\"{index:02d}_{name}\"\n        source = mapping[name].resolve()\n        try:\n            link.symlink_to(source, target_is_directory=True)\n        except (OSError, NotImplementedError):\n            shutil.copytree(source, link)\n    return destination\n\n\ndef organise_from_metadata_csv(csv_path: Path, image_dirs: list[Path], id_column: str,\n                               label_column: str, class_names: list[str], destination: Path,\n                               label_map: dict | None = None,\n                               id_suffixes: tuple[str, ...] = (\".jpg\", \".jpeg\", \".png\"),\n                               row_filter: tuple[str, str] | None = None) -> dict:\n    \"\"\"Turn 'CSV of labels + flat folder of images' into folder-per-class.\n\n    Two of the three image datasets ship this way rather than as class folders —\n    APTOS gives you train.csv with (id_code, diagnosis) and one flat directory,\n    HAM10000 gives you metadata with (image_id, dx) and images split across two.\n    Returns {class_name: directory}, or {} if the CSV or the named columns are\n    missing, in which case the caller falls back as usual.\n\n    Columns are looked up by the constants at the top of each disease script, so a\n    mismatch is a one-line fix in an obvious place rather than a hunt through here.\n    \"\"\"\n    import pandas as pd\n\n    if not csv_path.exists():\n        print(f\"[data] no metadata CSV at {csv_path}\")\n        return {}\n\n    frame = pd.read_csv(csv_path)\n    if row_filter:\n        column, value = row_filter\n        if column not in frame.columns:\n            print(f\"[warn] {csv_path.name} has no column {column!r} to filter on. \"\n                  f\"Available: {list(frame.columns)}\")\n            return {}\n        before = len(frame)\n        frame = frame[frame[column].astype(str) == str(value)]\n        print(f\"[data] {column}=={value!r}: {len(frame)} of {before} rows\")\n\n    missing = [c for c in (id_column, label_column) if c not in frame.columns]\n    if missing:\n        print(f\"[warn] {csv_path.name} has no column(s) {missing}. \"\n              f\"Available: {list(frame.columns)}\")\n        print(\"[warn] fix ID_COLUMN / LABEL_COLUMN at the top of this script.\")\n        return {}\n\n    index: dict[str, Path] = {}\n    macos_debris = 0\n    for directory in image_dirs:\n        if not directory.exists():\n            continue\n        for path in directory.rglob(\"*\"):\n            if path.is_file() and path.suffix.lower() in id_suffixes:\n                if is_macos_metadata(path):\n                    macos_debris += 1\n                    continue\n                index.setdefault(path.stem, path)\n    report_macos_debris(macos_debris, str(csv_path.parent))\n\n    if not index:\n        print(f\"[warn] no images found under {[str(d) for d in image_dirs]}\")\n        return {}\n\n    if destination.exists():\n        shutil.rmtree(destination)\n\n    linked = {name: 0 for name in class_names}\n    skipped_labels: dict[str, int] = {}\n    missing_images = 0\n    for _, row in frame.iterrows():\n        raw = row[label_column]\n        label = label_map.get(raw, label_map.get(str(raw))) if label_map else str(raw)\n        if label not in linked:\n            # A label outside CLASS_NAMES is dropped, but never silently — an\n            # unexpected category usually means the taxonomy moved under you.\n            skipped_labels[str(label)] = skipped_labels.get(str(label), 0) + 1\n            continue\n        source = index.get(str(row[id_column]))\n        if source is None:\n            missing_images += 1\n            continue\n        class_dir = destination / label\n        class_dir.mkdir(parents=True, exist_ok=True)\n        target = class_dir / source.name\n        if not target.exists():\n            try:\n                target.symlink_to(source.resolve())\n            except (OSError, NotImplementedError):\n                shutil.copy2(source, target)\n        linked[label] += 1\n\n    found = {name: destination / name for name, count in linked.items() if count > 0}\n    print(f\"[data] organised from {csv_path.name}: \"\n          + \", \".join(f\"{n}={linked[n]}\" for n in class_names))\n    if skipped_labels:\n        print(f\"[data] rows dropped, label not in CLASS_NAMES: {skipped_labels}\")\n    if missing_images:\n        print(f\"[warn] {missing_images} rows had no matching image file. \"\n              f\"If this is most of them, check ID_COLUMN against the filenames.\")\n    return found\n\n\n# --------------------------------------------------------------------------\n# Data pipeline\n# --------------------------------------------------------------------------\n\nclass RetinaCrop:\n    \"\"\"Crop a fundus photograph to its illuminated circle, dropping the black frame.\n\n    Retinal cameras put a circle of retina inside a large black rectangle, and on\n    APTOS that border is often a third of the frame. A plain resize spends its 224\n    pixels on the border as well as the retina, so the only part of the image\n    carrying signal arrives at a fraction of the resolution it could have. Cropping\n    first is the cheap half of Ben Graham's preprocessing and most of where the\n    benefit comes from.\n\n    Deliberately conservative: if the detected region is implausibly small the\n    original image is returned untouched. A crop that silently ate the retina would\n    be far worse than no crop, and it would be invisible in the metrics.\n    \"\"\"\n\n    def __init__(self, threshold: int = 10, min_area_fraction: float = 0.10):\n        self.threshold = threshold\n        self.min_area_fraction = min_area_fraction\n\n    def __call__(self, image: Image.Image) -> Image.Image:\n        array = np.asarray(image.convert(\"L\"))\n        mask = array > self.threshold\n        if not mask.any():\n            return image\n        rows = np.where(mask.any(axis=1))[0]\n        cols = np.where(mask.any(axis=0))[0]\n        top, bottom, left, right = int(rows[0]), int(rows[-1]), int(cols[0]), int(cols[-1])\n        height, width = bottom - top + 1, right - left + 1\n        if height * width < self.min_area_fraction * array.size:\n            return image\n        return image.crop((left, top, right + 1, bottom + 1))\n\n\ndef build_transforms(image_size: int, retina_crop: bool = False,\n                     crop_threshold: int = 10):\n    pre = [RetinaCrop(crop_threshold)] if retina_crop else []\n    train_tf = transforms.Compose(pre + [\n        transforms.Resize((image_size, image_size)),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomRotation(10),\n        transforms.ColorJitter(brightness=0.15, contrast=0.15),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    eval_tf = transforms.Compose(pre + [\n        transforms.Resize((image_size, image_size)),\n        transforms.ToTensor(),\n        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n    return train_tf, eval_tf\n\n\ndef _cap_per_class(dataset: datasets.ImageFolder, cap: int | None, seed: int) -> list[int]:\n    indices = list(range(len(dataset)))\n    if cap is None:\n        return indices\n    rng = random.Random(seed)\n    by_class: dict[int, list[int]] = {}\n    for i, (_, label) in enumerate(dataset.samples):\n        by_class.setdefault(label, []).append(i)\n    kept: list[int] = []\n    for label, items in by_class.items():\n        rng.shuffle(items)\n        kept.extend(items[:cap])\n    kept.sort()\n    return kept\n\n\ndef stratified_split(dataset: datasets.ImageFolder, indices: list[int],\n                     val_split: float, seed: int) -> tuple[list[int], list[int]]:\n    \"\"\"Split per class, so a rare class cannot land entirely in one side.\n\n    Severe DR and dermatofibroma are single-digit percentages of their datasets;\n    a random split can leave validation with zero examples of them, which reports\n    a flattering accuracy for a model that never learned them at all.\n    \"\"\"\n    rng = random.Random(seed)\n    by_class: dict[int, list[int]] = {}\n    for i in indices:\n        by_class.setdefault(dataset.samples[i][1], []).append(i)\n\n    train_idx: list[int] = []\n    val_idx: list[int] = []\n    for label, items in sorted(by_class.items()):\n        rng.shuffle(items)\n        n_val = max(1, int(round(len(items) * val_split))) if len(items) > 1 else 0\n        val_idx.extend(items[:n_val])\n        train_idx.extend(items[n_val:])\n    return sorted(train_idx), sorted(val_idx)\n\n\nclass ListImageDataset(torch.utils.data.Dataset):\n    \"\"\"Images from an explicit (path, label) list.\n\n    Used when the validation set is supplied separately rather than split out of\n    training. ImageFolder cannot serve that case safely: it derives labels from\n    sorted directory names *per root*, so a validation set missing one class would\n    renumber every label after it, and the model would be scored against a\n    silently shifted key. Carrying the label with the path removes the failure.\n    \"\"\"\n\n    def __init__(self, items: list[tuple[str, int]], transform):\n        self.items = items\n        self.transform = transform\n\n    def __len__(self) -> int:\n        return len(self.items)\n\n    def __getitem__(self, index: int):\n        path, label = self.items[index]\n        with Image.open(path) as image:\n            return self.transform(image.convert(\"RGB\")), label\n\n\ndef items_from_mapping(mapping: dict[str, Path], class_names: list[str],\n                       cap: int | None = None, seed: int = 42) -> list[tuple[str, int]]:\n    \"\"\"Flatten {class: directory} into (path, label) pairs indexed by class_names.\n\n    The label is the position in class_names, which is the same order labels.json\n    publishes — so train and validation stay on one numbering even when they come\n    from different datasets entirely.\n    \"\"\"\n    rng = random.Random(seed)\n    items: list[tuple[str, int]] = []\n    macos_debris = 0\n    for label, name in enumerate(class_names):\n        directory = mapping.get(name)\n        if directory is None:\n            continue\n        candidates = [p for p in Path(directory).rglob(\"*\")\n                      if p.is_file() and p.suffix.lower() in IMAGE_SUFFIXES]\n        files = sorted(p for p in candidates if not is_macos_metadata(p))\n        macos_debris += len(candidates) - len(files)\n        if cap is not None and len(files) > cap:\n            rng.shuffle(files)\n            files = files[:cap]\n        items.extend((str(p), label) for p in files)\n    report_macos_debris(macos_debris)\n    return items\n\n\ndef split_items_stratified(items: list[tuple[str, int]], fraction: float,\n                           seed: int) -> tuple[list[tuple[str, int]], list[tuple[str, int]]]:\n    \"\"\"Carve a stratified slice off an item list. Returns (remainder, slice).\n\n    Stratified per class for the usual reason: on an imbalanced set a random slice\n    can take every example of a rare class, or none of them, and a final test\n    number computed on a slice missing a class is not measuring what it claims.\n    Never takes an entire class, so the remainder always keeps something to score.\n    \"\"\"\n    if fraction <= 0 or not items:\n        return items, []\n\n    rng = random.Random(seed)\n    by_label: dict[int, list[tuple[str, int]]] = {}\n    for item in items:\n        by_label.setdefault(item[1], []).append(item)\n\n    remainder: list[tuple[str, int]] = []\n    held_out: list[tuple[str, int]] = []\n    for label, group in sorted(by_label.items()):\n        group = sorted(group)          # deterministic order before shuffling\n        rng.shuffle(group)\n        if len(group) > 1:\n            take = min(max(int(round(len(group) * fraction)), 1), len(group) - 1)\n        else:\n            take = 0                   # one example of a class stays measurable\n        held_out.extend(group[:take])\n        remainder.extend(group[take:])\n    return sorted(remainder), sorted(held_out)\n\n\ndef balanced_sampler_from_labels(labels: list[int]) -> WeightedRandomSampler:\n    counts: dict[int, int] = {}\n    for label in labels:\n        counts[label] = counts.get(label, 0) + 1\n    weights = [1.0 / counts[label] for label in labels]\n    return WeightedRandomSampler(weights, num_samples=len(labels), replacement=True)\n\n\ndef make_balanced_sampler(dataset: datasets.ImageFolder, indices: list[int]) -> WeightedRandomSampler:\n    \"\"\"Sample rare classes up, so the model cannot score well by ignoring them.\n\n    Every Phase 1 image dataset is heavily imbalanced — roughly 70% of the DR set\n    is No_DR, and 'nv' is about two thirds of HAM10000. A model that predicts the\n    majority class for everything looks fine on accuracy while being useless for\n    screening, which is the one thing this product must not ship.\n    \"\"\"\n    counts: dict[int, int] = {}\n    for i in indices:\n        label = dataset.samples[i][1]\n        counts[label] = counts.get(label, 0) + 1\n    weights = [1.0 / counts[dataset.samples[i][1]] for i in indices]\n    return WeightedRandomSampler(weights, num_samples=len(indices), replacement=True)\n\n\n# --------------------------------------------------------------------------\n# Model\n# --------------------------------------------------------------------------\n\ndef build_model(backbone: str, num_classes: int) -> tuple[nn.Module, bool]:\n    \"\"\"Return (model, pretrained). Falls back to random init when offline.\n\n    Kaggle notebooks run without internet unless it is switched on, so a bare\n    `weights=DEFAULT` is a coin flip on whether the script survives its first\n    line. Try for pretrained, carry on without it, and record which happened in\n    metrics.json — the difference matters a great deal when reading the scores.\n    \"\"\"\n    if backbone not in SUPPORTED_BACKBONES:\n        raise ValueError(f\"backbone must be one of {SUPPORTED_BACKBONES}, got {backbone!r}\")\n\n    pretrained = True\n    try:\n        if backbone == \"efficientnet_b0\":\n            model = models.efficientnet_b0(weights=models.EfficientNet_B0_Weights.DEFAULT)\n        else:\n            model = models.mobilenet_v2(weights=models.MobileNet_V2_Weights.DEFAULT)\n    except Exception as exc:  # no internet, or no cached weights\n        print(f\"[warn] pretrained weights unavailable ({type(exc).__name__}); \"\n              f\"training {backbone} from scratch. Expect much weaker scores.\")\n        pretrained = False\n        model = (models.efficientnet_b0(weights=None) if backbone == \"efficientnet_b0\"\n                 else models.mobilenet_v2(weights=None))\n\n    in_features = model.classifier[-1].in_features\n    model.classifier[-1] = nn.Linear(in_features, num_classes)\n    return model, pretrained\n\n\n# --------------------------------------------------------------------------\n# Metrics\n# --------------------------------------------------------------------------\n\ndef confusion_matrix(true: np.ndarray, pred: np.ndarray, num_classes: int) -> np.ndarray:\n    matrix = np.zeros((num_classes, num_classes), dtype=int)\n    for t, p in zip(true, pred):\n        matrix[int(t), int(p)] += 1\n    return matrix\n\n\ndef score(true: np.ndarray, pred: np.ndarray, class_names: list[str]) -> dict:\n    \"\"\"Accuracy, macro F1, and per-class precision/recall.\n\n    Recall is the number to read first. The backend turns a positive into\n    `needs_professional_review` — a human looks at it — so a false positive costs\n    someone's attention while a false negative sends a sick patient home\n    reassured. The two are not equally bad, and accuracy hides the difference.\n    \"\"\"\n    n = len(class_names)\n    matrix = confusion_matrix(true, pred, n)\n    support = matrix.sum(axis=1)\n    predicted = matrix.sum(axis=0)\n    correct = np.diag(matrix)\n\n    with np.errstate(divide=\"ignore\", invalid=\"ignore\"):\n        recall = np.where(support > 0, correct / np.maximum(support, 1), 0.0)\n        precision = np.where(predicted > 0, correct / np.maximum(predicted, 1), 0.0)\n        f1 = np.where((precision + recall) > 0,\n                      2 * precision * recall / np.maximum(precision + recall, 1e-12), 0.0)\n\n    return {\n        \"accuracy\": float(correct.sum() / max(len(true), 1)),\n        \"macro_f1\": float(f1.mean()),\n        \"macro_recall\": float(recall.mean()),\n        \"per_class\": {\n            name: {\n                \"precision\": round(float(precision[i]), 4),\n                \"recall\": round(float(recall[i]), 4),\n                \"f1\": round(float(f1[i]), 4),\n                \"support\": int(support[i]),\n            }\n            for i, name in enumerate(class_names)\n        },\n        \"confusion_matrix\": matrix.tolist(),\n        \"confusion_matrix_axes\": {\"rows\": \"true\", \"columns\": \"predicted\", \"order\": class_names},\n    }\n\n\n# --------------------------------------------------------------------------\n# Train / evaluate\n# --------------------------------------------------------------------------\n\ndef run_epoch(model, loader, device, criterion, optimiser=None) -> tuple[float, np.ndarray, np.ndarray]:\n    training = optimiser is not None\n    model.train(training)\n    total_loss, seen = 0.0, 0\n    trues, preds = [], []\n\n    with torch.set_grad_enabled(training):\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            if training:\n                optimiser.zero_grad()\n                loss.backward()\n                optimiser.step()\n            total_loss += float(loss.item()) * labels.size(0)\n            seen += labels.size(0)\n            trues.append(labels.detach().cpu().numpy())\n            preds.append(outputs.detach().argmax(1).cpu().numpy())\n\n    return (total_loss / max(seen, 1),\n            np.concatenate(trues) if trues else np.array([]),\n            np.concatenate(preds) if preds else np.array([]))\n\n\ndef set_seed(seed: int) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\n# --------------------------------------------------------------------------\n# Artifacts\n# --------------------------------------------------------------------------\n\ndef write_artifacts(output_dir: Path, config: TrainingConfig, model: nn.Module,\n                    class_names: list[str], metrics: dict, extra: dict) -> None:\n    output_dir.mkdir(parents=True, exist_ok=True)\n\n    torch.save(\n        {\n            \"state_dict\": model.state_dict(),\n            \"backbone\": config.backbone,\n            \"num_classes\": len(class_names),\n            \"class_names\": class_names,\n            \"disease_id\": config.disease_id,\n            \"model_version\": config.model_version,\n        },\n        output_dir / \"model.pt\",\n    )\n\n    (output_dir / \"labels.json\").write_text(json.dumps({\n        \"disease_id\": config.disease_id,\n        \"class_names\": class_names,\n        \"index_to_label\": {str(i): name for i, name in enumerate(class_names)},\n    }, indent=2), encoding=\"utf-8\")\n\n    (output_dir / \"metrics.json\").write_text(json.dumps({\n        \"disease_id\": config.disease_id,\n        \"backbone\": config.backbone,\n        \"model_version\": config.model_version,\n        \"is_demo_model\": True,\n        **extra,\n        **metrics,\n        \"config\": asdict(config),\n        \"environment\": {\n            \"python\": platform.python_version(),\n            \"torch\": torch.__version__,\n            \"device\": extra.get(\"device\", \"unknown\"),\n        },\n    }, indent=2), encoding=\"utf-8\")\n\n    # The backend must reproduce this transform exactly or the model sees a\n    # different distribution than it was trained on.\n    (output_dir / \"preprocessing.json\").write_text(json.dumps({\n        \"disease_id\": config.disease_id,\n        \"input\": \"rgb_image\",\n        \"resize\": [config.image_size, config.image_size],\n        \"interpolation\": \"bilinear\",\n        \"scale\": \"divide_by_255\",\n        \"normalize\": {\"mean\": IMAGENET_MEAN, \"std\": IMAGENET_STD},\n        \"channel_order\": \"RGB\",\n        \"tensor_layout\": \"NCHW\",\n        # The backend has to apply this in the same order, crop included. A model\n        # trained on cropped retinas and served uncropped ones sees a different\n        # distribution and says so in neither the checkpoint nor its output.\n        \"retina_crop\": {\n            \"enabled\": config.retina_crop,\n            \"luminance_threshold\": config.retina_crop_threshold,\n            \"method\": \"bounding box of pixels brighter than the threshold, applied \"\n                      \"before resize; skipped when that box is under 10% of the frame\",\n        },\n        \"notes\": (\"Apply retina_crop (if enabled), then resize, then to-tensor (0-1), \"\n                  \"then normalize. No centre crop.\"),\n    }, indent=2), encoding=\"utf-8\")\n\n\n# --------------------------------------------------------------------------\n# Entry point\n# --------------------------------------------------------------------------\n\ndef train(config: TrainingConfig, class_dirs: dict | None = None,\n          val_class_dirs: dict | None = None,\n          data_source_label: str | None = None) -> dict:\n    \"\"\"Train one disease/backbone pair and write the four artifacts.\n\n    `class_dirs` lets a disease script hand over a {class: directory} mapping it\n    built itself — from a metadata CSV, say. When it is None the dataset folder is\n    searched instead. Either way a failure degrades to synthetic data rather than\n    raising, so the script always finishes and always says what it actually ran on.\n\n    `val_class_dirs` supplies the validation set explicitly instead of splitting it\n    out of training. Two Phase 1 diseases need this and for the same underlying\n    reason — a split someone else already made is better than one we invent. Skin\n    ships a `split` column that keeps a patient's images on one side, which a random\n    split would break by putting two photos of one lesion in train and test. DR\n    validates on IDRiD, a different cohort and camera entirely, which is a far\n    honest test of generalisation than a held-out slice of the training set.\n    When supplied, nothing from it ever reaches training.\n\n    `config.final_test_fraction` carves a slice off that holdout before the first\n    epoch and scores it exactly once, after the weights are written, into\n    final_test_metrics.json. The remainder stays the epoch-by-epoch validation\n    signal. The separation is the whole point: a number you watch while training is\n    a number you have implicitly tuned against, so it flatters the model. This one\n    is read after the fact and cannot be steered by.\n    \"\"\"\n    started = time.time()\n    set_seed(config.seed)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    output_dir = config.resolved_output_dir()\n    print(f\"[{config.disease_id}/{config.backbone}] device={device}\")\n\n    if class_dirs:\n        root, mapping = Path(config.dataset_dir), class_dirs\n        # Default suits the CSV-driven diseases; folder-driven callers say so.\n        data_source = data_source_label or \"kaggle-input-via-metadata-csv\"\n    else:\n        root, mapping = find_class_root(Path(config.dataset_dir), config.class_names)\n        data_source = \"kaggle-input\"\n    if not mapping:\n        root = None\n\n    if root is None:\n        print(f\"[warn] no class folders found under {config.dataset_dir!r}.\")\n        print(\"[warn] falling back to a synthetic dataset so the pipeline still runs.\")\n        print(\"[warn] fix DATASET_DIR / CLASS_NAMES at the top of the disease script \"\n              \"once the real dataset is attached.\")\n        root = build_synthetic_dataset(\n            working_root() / f\"_synthetic_{config.disease_id}\",\n            config.class_names, config.image_size)\n        data_source = \"synthetic-fallback\"\n        mapping = {name: root / name for name in config.class_names}\n\n    print(f\"[data] root={root}\")\n    for name in config.class_names:\n        found = mapping.get(name)\n        print(f\"[data]   {name}: {'-> ' + found.name if found else 'MISSING'}\")\n\n    train_tf, eval_tf = build_transforms(config.image_size, config.retina_crop,\n                                         config.retina_crop_threshold)\n    found_names = [n for n in config.class_names if n in mapping]\n\n    if val_class_dirs:\n        train_items = items_from_mapping(mapping, found_names,\n                                         config.max_images_per_class, config.seed)\n        holdout_items = items_from_mapping(val_class_dirs, found_names, None, config.seed)\n        # Split before anything is trained. Everything downstream only ever sees\n        # `val_items`; `final_test_items` is not referenced again until after the\n        # loop, which is what makes its number worth reading.\n        val_items, final_test_items = split_items_stratified(\n            holdout_items, config.final_test_fraction, config.seed + 1)\n        if not train_items:\n            raise RuntimeError(f\"no training images found under {root}; check DATASET_DIR\")\n        if not val_items:\n            print(\"[warn] the supplied validation set contained no images; \"\n                  \"scoring against the training set instead. Check its constants.\")\n            val_items = train_items\n        validation_source = \"provided-holdout\"\n        n_train, n_val = len(train_items), len(val_items)\n        print(f\"[data] {n_train} train / {n_val} val (validation supplied separately, \"\n              f\"never trained on) across {len(found_names)} classes\")\n        if final_test_items:\n            print(f\"[data] {len(final_test_items)} images sequestered as a FINAL TEST \"\n                  f\"slice — not looked at until training has finished\")\n\n        train_loader = DataLoader(ListImageDataset(train_items, train_tf),\n                                  batch_size=config.batch_size,\n                                  sampler=balanced_sampler_from_labels([l for _, l in train_items]),\n                                  num_workers=config.num_workers, drop_last=False)\n        val_loader = DataLoader(ListImageDataset(val_items, eval_tf),\n                                batch_size=config.batch_size, shuffle=False,\n                                num_workers=config.num_workers)\n    else:\n        view = build_class_view(mapping, config.class_names,\n                                working_root() / f\"_view_{config.disease_id}_{config.backbone}\")\n\n        # Counted rather than merely rejected: ImageFolder would otherwise drop\n        # these without a word, and the next unreadable file would look the same.\n        macos_debris = {\"n\": 0}\n\n        def _is_image(path: str) -> bool:\n            candidate = Path(path)\n            if candidate.suffix.lower() not in IMAGE_SUFFIXES:\n                return False\n            if is_macos_metadata(candidate):\n                macos_debris[\"n\"] += 1\n                return False\n            return True\n\n        base = datasets.ImageFolder(str(view), transform=train_tf, is_valid_file=_is_image)\n        report_macos_debris(macos_debris[\"n\"], str(view))\n        # eval_base re-walks the identical tree; zero it so the count is not doubled.\n        macos_debris[\"n\"] = 0\n        eval_base = datasets.ImageFolder(str(view), transform=eval_tf, is_valid_file=_is_image)\n        base.classes = eval_base.classes = found_names\n\n        if not base.samples:\n            raise RuntimeError(f\"no images found under {root}; check DATASET_DIR\")\n\n        final_test_items = []\n        if config.final_test_fraction > 0:\n            print(\"[warn] final_test_fraction is set but no holdout was supplied, so \"\n                  \"there is nothing to carve a final test slice from. Skipping it — \"\n                  \"a slice taken from the same split used for validation would not \"\n                  \"be independent of it.\")\n        indices = _cap_per_class(base, config.max_images_per_class, config.seed)\n        train_idx, val_idx = stratified_split(base, indices, config.val_split, config.seed)\n        validation_source = \"internal-stratified-split\"\n        n_train, n_val = len(train_idx), len(val_idx)\n        print(f\"[data] {n_train} train / {n_val} val across {len(found_names)} classes\")\n\n        train_loader = DataLoader(Subset(base, train_idx), batch_size=config.batch_size,\n                                  sampler=make_balanced_sampler(base, train_idx),\n                                  num_workers=config.num_workers, drop_last=False)\n        val_loader = DataLoader(Subset(eval_base, val_idx or train_idx),\n                                batch_size=config.batch_size, shuffle=False,\n                                num_workers=config.num_workers)\n\n    model, pretrained = build_model(config.backbone, len(found_names))\n    model.to(device)\n    criterion = nn.CrossEntropyLoss()\n    optimiser = torch.optim.AdamW(model.parameters(), lr=config.learning_rate)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimiser, T_max=max(config.epochs, 1))\n\n    history = []\n    for epoch in range(1, config.epochs + 1):\n        train_loss, t_true, t_pred = run_epoch(model, train_loader, device, criterion, optimiser)\n        val_loss, v_true, v_pred = run_epoch(model, val_loader, device, criterion)\n        scheduler.step()\n        train_acc = float((t_true == t_pred).mean()) if len(t_true) else 0.0\n        val_acc = float((v_true == v_pred).mean()) if len(v_true) else 0.0\n        history.append({\"epoch\": epoch, \"train_loss\": round(train_loss, 4),\n                        \"val_loss\": round(val_loss, 4),\n                        \"train_accuracy\": round(train_acc, 4),\n                        \"val_accuracy\": round(val_acc, 4)})\n        print(f\"[epoch {epoch}/{config.epochs}] train_loss={train_loss:.4f} \"\n              f\"val_loss={val_loss:.4f} val_acc={val_acc:.4f}\")\n\n    _, v_true, v_pred = run_epoch(model, val_loader, device, criterion)\n    metrics = score(v_true, v_pred, found_names)\n\n    write_artifacts(output_dir, config, model, found_names, metrics, {\n        \"data_source\": data_source,\n        \"dataset_root\": str(root),\n        \"pretrained_backbone\": pretrained,\n        \"device\": str(device),\n        \"train_size\": n_train,\n        \"val_size\": n_val,\n        \"validation_source\": validation_source,\n        \"final_test_size\": len(final_test_items),\n        \"final_test_influenced_training\": False,\n        \"history\": history,\n        \"duration_seconds\": round(time.time() - started, 1),\n    })\n\n    # ---- The one and only look at the final test slice. ----\n    # Deliberately placed after write_artifacts: the weights are already on disk,\n    # so nothing computed here can loop back into training, checkpointing or\n    # hyperparameters. Reading it earlier would make it just another validation set.\n    artifact_names = [\"model.pt\", \"labels.json\", \"metrics.json\", \"preprocessing.json\"]\n    if final_test_items:\n        final_loader = DataLoader(ListImageDataset(final_test_items, eval_tf),\n                                  batch_size=config.batch_size, shuffle=False,\n                                  num_workers=config.num_workers)\n        _, f_true, f_pred = run_epoch(model, final_loader, device, criterion)\n        final_metrics = score(f_true, f_pred, found_names)\n        (output_dir / \"final_test_metrics.json\").write_text(json.dumps({\n            \"disease_id\": config.disease_id,\n            \"backbone\": config.backbone,\n            \"model_version\": config.model_version,\n            \"is_demo_model\": True,\n            \"scored_once_after_training\": True,\n            \"used_for_any_training_decision\": False,\n            \"n_images\": len(final_test_items),\n            \"fraction_of_supplied_holdout\": config.final_test_fraction,\n            \"split_seed\": config.seed + 1,\n            \"note\": (\"Held out before the first epoch and scored once, after the \"\n                     \"weights above were written. metrics.json is the validation \"\n                     \"number watched during training; this is the one that was \"\n                     \"never optimised against.\"),\n            **final_metrics,\n        }, indent=2), encoding=\"utf-8\")\n        artifact_names.append(\"final_test_metrics.json\")\n\n    print(f\"\\n[done] accuracy={metrics['accuracy']:.4f} macro_f1={metrics['macro_f1']:.4f} \"\n          f\"macro_recall={metrics['macro_recall']:.4f}   (validation, watched each epoch)\")\n    if final_test_items:\n        print(f\"[done] FINAL TEST accuracy={final_metrics['accuracy']:.4f} \"\n              f\"macro_f1={final_metrics['macro_f1']:.4f} \"\n              f\"macro_recall={final_metrics['macro_recall']:.4f}   \"\n              f\"(scored once, on {len(final_test_items)} unseen images)\")\n    print(f\"[done] artifacts -> {output_dir}\")\n    for name in artifact_names:\n        print(f\"[done]   {name}\")\n    if data_source == \"synthetic-fallback\":\n        print(\"[done] NOTE: trained on synthetic data — scores are meaningless. \"\n              \"Attach the real dataset and set DATASET_DIR.\")\n    return metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T19:38:13.042025Z","iopub.execute_input":"2026-08-30T19:38:13.042712Z","iopub.status.idle":"2026-08-30T19:38:13.141313Z","shell.execute_reply.started":"2026-08-30T19:38:13.042679Z","shell.execute_reply":"2026-08-30T19:38:13.140490Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# APTOS / DIABETIC RETINOPATHY TRAINING\n# FINAL CUDA-SAFE VERSION\n# ============================================================\n\n# IMPORTANT:\n# Restart the Kaggle session before running this cell if you\n# previously got a CUDA error.\n\nimport os\n\n# Must be set before any attempt to use CUDA in this fresh session\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"\"\n\nfrom pathlib import Path\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\n\n\n# ============================================================\n# FORCE THE FIRST-CELL TRAINING PIPELINE TO USE CPU\n# ============================================================\n\nCPU_DEVICE = torch.device(\"cpu\")\n\n\ndef force_cpu_device(*args, **kwargs):\n    \"\"\"Always return CPU.\"\"\"\n    return CPU_DEVICE\n\n\n# The important part:\n# train() keeps its own globals dictionary from the first cell.\n# We patch THAT dictionary, not just this cell's globals.\n\nif \"train\" not in globals():\n    raise RuntimeError(\n        \"The first cnn_training cell has not been run. \"\n        \"Run Cell 1 first, then run this cell.\"\n    )\n\n\nTRAIN_GLOBALS = train.__globals__\n\n# Patch common device-selector names if they exist\nfor name in [\n    \"get_device\",\n    \"_get_device\",\n    \"select_device\",\n    \"choose_device\",\n    \"device\",\n]:\n    if name in TRAIN_GLOBALS:\n        TRAIN_GLOBALS[name] = force_cpu_device\n\n\n# Also force PyTorch's CUDA availability check to return False.\n# This catches training code using:\n# torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n_original_cuda_available = torch.cuda.is_available\ntorch.cuda.is_available = lambda: False\n\n\nprint(\"=\" * 75)\nprint(\"DEVICE SAFETY CHECK\")\nprint(\"=\" * 75)\nprint(\"PyTorch device:\", CPU_DEVICE)\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"✓ CPU mode forced inside the training pipeline\")\nprint(\"=\" * 75)\n\n\n# ============================================================\n# RANDOM SEED\n# ============================================================\n\nSEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\n\n# ============================================================\n# FIND THE APTOS DATASET\n# ============================================================\n\nINPUT_ROOT = Path(\"/kaggle/input\")\n\nif not INPUT_ROOT.exists():\n    raise FileNotFoundError(\"/kaggle/input does not exist.\")\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"SEARCHING FOR DIABETIC RETINOPATHY DATASET\")\nprint(\"=\" * 75)\n\nprint(\"\\nAttached datasets:\")\nfor item in INPUT_ROOT.iterdir():\n    print(\"📁\", item.name)\n\n\n# ============================================================\n# FIND APTOS-STYLE CSV\n# ============================================================\n\ncsv_candidates = []\n\nfor csv_file in INPUT_ROOT.rglob(\"*.csv\"):\n    try:\n        sample = pd.read_csv(csv_file, nrows=10)\n        cols = {\n            str(c).strip().lower()\n            for c in sample.columns\n        }\n\n        if (\n            \"diagnosis\" in cols\n            and any(\n                x in cols\n                for x in [\"id_code\", \"image_id\", \"image\", \"id\"]\n            )\n        ):\n            csv_candidates.append(csv_file)\n\n    except Exception:\n        continue\n\n\nif not csv_candidates:\n    print(\"\\nCSV files found:\")\n    for file in INPUT_ROOT.rglob(\"*.csv\"):\n        print(\" -\", file)\n\n    raise FileNotFoundError(\n        \"\\nCould not find a CSV with image IDs and diagnosis labels.\"\n    )\n\n\n# Prefer train.csv\ncsv_candidates.sort(\n    key=lambda p: (\n        \"train\" not in p.name.lower(),\n        len(str(p))\n    )\n)\n\nMETADATA_PATH = csv_candidates[0]\nDATASET_DIR = METADATA_PATH.parent\n\nprint(\"\\n✓ Using metadata CSV:\")\nprint(METADATA_PATH)\n\n\n# ============================================================\n# READ AND IDENTIFY COLUMNS\n# ============================================================\n\ndf = pd.read_csv(METADATA_PATH)\n\nprint(\"\\nColumns:\", list(df.columns))\nprint(\"Total CSV rows:\", len(df))\n\ncolumn_lookup = {\n    str(col).strip().lower(): col\n    for col in df.columns\n}\n\n\nID_COLUMN = None\nfor name in [\"id_code\", \"image_id\", \"image\", \"id\"]:\n    if name in column_lookup:\n        ID_COLUMN = column_lookup[name]\n        break\n\n\nLABEL_COLUMN = None\nfor name in [\"diagnosis\", \"grade\", \"severity\", \"label\", \"dr_grade\"]:\n    if name in column_lookup:\n        LABEL_COLUMN = column_lookup[name]\n        break\n\n\nif ID_COLUMN is None or LABEL_COLUMN is None:\n    raise ValueError(\n        f\"Could not identify required columns.\\n\"\n        f\"Available columns: {list(df.columns)}\"\n    )\n\n\nprint(f\"\\n✓ Image ID column: {ID_COLUMN}\")\nprint(f\"✓ Label column: {LABEL_COLUMN}\")\n\n\n# ============================================================\n# STANDARD DIABETIC RETINOPATHY CLASSES\n# ============================================================\n\nCLASS_NAMES = [\"0\", \"1\", \"2\", \"3\", \"4\"]\n\n# Convert labels to strings so they match CLASS_NAMES\ndf[LABEL_COLUMN] = df[LABEL_COLUMN].astype(str).str.strip()\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"CLASS DISTRIBUTION\")\nprint(\"=\" * 75)\nprint(df[LABEL_COLUMN].value_counts().sort_index())\n\n\n# Save cleaned CSV because organiser reads directly from disk\nCLEANED_CSV = Path(\"/kaggle/working/aptos_cleaned_labels.csv\")\ndf.to_csv(CLEANED_CSV, index=False)\n\n\n# ============================================================\n# FIND IMAGE DIRECTORIES\n# ============================================================\n\nIMAGE_DIRS = []\n\n# Search the dataset and all its subdirectories for image folders\nfor directory in [DATASET_DIR] + [\n    p for p in DATASET_DIR.rglob(\"*\") if p.is_dir()\n]:\n    try:\n        if any(\n            f.is_file() and f.suffix.lower() in IMAGE_SUFFIXES\n            for f in directory.iterdir()\n        ):\n            IMAGE_DIRS.append(directory)\n    except (PermissionError, OSError):\n        pass\n\n\nIMAGE_DIRS = list(dict.fromkeys(IMAGE_DIRS))\n\nif not IMAGE_DIRS:\n    raise FileNotFoundError(\n        \"No folders containing images were found.\"\n    )\n\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"IMAGE FOLDERS FOUND\")\nprint(\"=\" * 75)\n\nfor directory in IMAGE_DIRS:\n    print(\"✓\", directory)\n\n\n# ============================================================\n# ORGANISE IMAGES INTO THE 5 CLASSES\n# ============================================================\n\nORGANIZED_DIR = (\n    working_root()\n    / \"_organised_diabetic_retinopathy\"\n)\n\nLABEL_MAP = {\n    \"0\": \"0\",\n    \"1\": \"1\",\n    \"2\": \"2\",\n    \"3\": \"3\",\n    \"4\": \"4\",\n    0: \"0\",\n    1: \"1\",\n    2: \"2\",\n    3: \"3\",\n    4: \"4\",\n}\n\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"ORGANISING DATASET\")\nprint(\"=\" * 75)\n\nclass_dirs = organise_from_metadata_csv(\n    csv_path=CLEANED_CSV,\n    image_dirs=IMAGE_DIRS,\n    id_column=ID_COLUMN,\n    label_column=LABEL_COLUMN,\n    class_names=CLASS_NAMES,\n    destination=ORGANIZED_DIR,\n    label_map=LABEL_MAP,\n)\n\n\nif not class_dirs:\n    raise RuntimeError(\n        \"No images matched the labels in the CSV.\"\n    )\n\n\nprint(\"\\n✓ Dataset organised successfully.\")\n\ntotal_images = 0\n\nfor class_name in CLASS_NAMES:\n    if class_name in class_dirs:\n        directory = Path(class_dirs[class_name])\n        count = len([\n            f for f in directory.iterdir()\n            if f.is_file()\n        ])\n        total_images += count\n        print(f\"Class {class_name}: {count} images\")\n\n\nprint(f\"\\nTotal images available: {total_images}\")\n\nif total_images == 0:\n    raise RuntimeError(\"Zero images were organised.\")\n\n\n# ============================================================\n# TRAINING CONFIGURATION\n# ============================================================\n\nconfig = TrainingConfig(\n    disease_id=\"diabetic_retinopathy\",\n    dataset_dir=str(DATASET_DIR),\n    class_names=CLASS_NAMES,\n    backbone=\"mobilenet_v2\",\n    image_size=224,\n\n    # Start small so we can verify everything works\n    epochs=6,\n    batch_size=16,\n    learning_rate=3e-4,\n    max_images_per_class=500,\n    num_workers=0,\n\n    output_dir=\"/kaggle/working/diabetic_retinopathy/mobilenet_v2\",\n)\n\n\n# ============================================================\n# EXTRA FINAL CPU PATCH\n# ============================================================\n\n# Re-patch immediately before calling train()\nTRAIN_GLOBALS = train.__globals__\n\nfor name in [\n    \"get_device\",\n    \"_get_device\",\n    \"select_device\",\n    \"choose_device\",\n]:\n    TRAIN_GLOBALS[name] = force_cpu_device\n\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"FINAL DEVICE CHECK BEFORE TRAINING\")\nprint(\"=\" * 75)\nprint(\"Forced device:\", force_cpu_device())\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"Starting training on CPU...\")\nprint(\"=\" * 75)\n\n\n# ============================================================\n# START TRAINING\n# ============================================================\n\nmetrics = train(\n    config=config,\n    class_dirs=class_dirs,\n)\n\n\n# ============================================================\n# RESULTS\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"🎉 TRAINING COMPLETED\")\nprint(\"=\" * 75)\n\nif isinstance(metrics, dict):\n    for key in [\"accuracy\", \"macro_f1\", \"macro_recall\"]:\n        if key in metrics:\n            print(f\"{key}: {metrics[key]:.4f}\")\n\n\nprint(\"\\nOutput folder:\")\nprint(config.resolved_output_dir())\n\nprint(\"\\nExpected artifacts:\")\nprint(\"✓ model.pt\")\nprint(\"✓ labels.json\")\nprint(\"✓ metrics.json\")\nprint(\"✓ preprocessing.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T19:38:13.142572Z","iopub.execute_input":"2026-08-30T19:38:13.142777Z","iopub.status.idle":"2026-08-30T20:14:34.870199Z","shell.execute_reply.started":"2026-08-30T19:38:13.142758Z","shell.execute_reply":"2026-08-30T20:14:34.869410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}