{"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\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            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() if d.is_dir()),\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    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                index.setdefault(path.stem, path)\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    for label, name in enumerate(class_names):\n        directory = mapping.get(name)\n        if directory is None:\n            continue\n        files = sorted(p for p in Path(directory).rglob(\"*\")\n                       if p.is_file() and p.suffix.lower() in IMAGE_SUFFIXES)\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    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        def _is_image(path: str) -> bool:\n            return Path(path).suffix.lower() in IMAGE_SUFFIXES\n\n        base = datasets.ImageFolder(str(view), transform=train_tf, is_valid_file=_is_image)\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-30T17:02:48.589677Z","iopub.execute_input":"2026-08-30T17:02:48.590017Z","iopub.status.idle":"2026-08-30T17:02:56.022704Z","shell.execute_reply.started":"2026-08-30T17:02:48.589976Z","shell.execute_reply":"2026-08-30T17:02:56.021844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"Diabetic retinopathy — EfficientNet-B0 training script.\n\nFive ordinal grades, and the set is dominated by No_DR — around seven in ten\nimages. Accuracy is therefore close to worthless here: a model that answers\n'No_DR' every time scores ~0.7 and would miss every patient who needs care.\nRead macro_recall and the per-class recall in metrics.json instead.\n\nTrains on APTOS, validates on IDRiD. The two datasets differ in more than name —\ndifferent cohort, different camera, different file extension, different label\nheader — so nothing is assumed to carry across; each has its own constants below.\n\nRun inside a Kaggle notebook:\n\n    !python ml/diseases/dr/train_efficientnet_b0.py\n\"\"\"\n\nimport pathlib\nimport sys\n\n# Two ways this script gets run, and it has to survive both.\n#\n# As a real file, __file__ exists: walk up to the repo root, put it on sys.path\n# and import the shared pipeline normally.\n#\n# Pasted into a Kaggle notebook, __file__ does not exist at all. The steps doc\n# has teammates paste ml/common/cnn_training.py into Cell 1 and this script into\n# Cell 2, so every name below is already defined in the notebook's globals and\n# there is nothing to import. Reading __file__ unguarded there raises NameError\n# before a single line of training config is read, so the access is guarded and\n# both paths fall through to whatever is already defined.\ntry:\n    _HERE = pathlib.Path(__file__).resolve()\nexcept NameError:          # pasted into a notebook cell; there is no file\n    _HERE = None\n\nif _HERE is not None:\n    for _candidate in _HERE.parents:\n        if (_candidate / \"ml\" / \"common\" / \"cnn_training.py\").exists():\n            sys.path.insert(0, str(_candidate))\n            break\n\n_NEEDED = (\"TrainingConfig\", \"organise_from_metadata_csv\", \"resolve_dataset_root\",\n           \"train\", \"working_root\")\ntry:\n    from ml.common.cnn_training import (TrainingConfig, organise_from_metadata_csv,\n                                        resolve_dataset_root, train, working_root)\nexcept ImportError as _exc:\n    # Cell 1 already defined these, so a failed import is expected in a notebook\n    # and harmless. Anything genuinely missing is a real problem, and saying so\n    # here beats a NameError several hundred lines further down.\n    _missing = [_n for _n in _NEEDED if _n not in globals()]\n    if _missing:\n        raise ImportError(\n            f\"{_missing} could not be imported from ml.common.cnn_training and \"\n            \"are not already defined. Running as a file? Check it still sits \"\n            \"inside the repo. In a notebook? Run the cnn_training.py cell first.\"\n        ) from _exc\n\n# ---------------------------------------------------------------------------\n# TRAINING DATA — APTOS  (CLAUDE.md > Datasets: aptos2019-blindness-detection)\n# ---------------------------------------------------------------------------\n# Two mounts, because Kaggle has two ways to attach the same competition data and\n# shows teammates whichever tab it feels like. The Datasets tab mounts at\n# /kaggle/input/<slug>; the Competitions tab inserts an extra segment and mounts\n# at /kaggle/input/competitions/<slug>. Hardcoding either one leaves the other\n# teammate with a silent fall-through to synthetic data that reads like a broken\n# path. Both are checked, in this order, and the run prints which it used.\n#\n# A third form turned up on a real session: /kaggle/input/datasets/<owner>/<slug>.\n# That one carries an owner segment, and for a self-uploaded dataset the owner is\n# whoever uploaded it — so it is globbed, never named. This is no longer specific\n# to APTOS (IDRiD mounts this way too), which is why the search itself now lives\n# in ml/common as resolve_dataset_root() and this script only supplies constants.\nDATASET_SLUG = \"aptos2019-blindness-detection\"\nDATASET_DIR_CANDIDATES = [\n    f\"/kaggle/input/{DATASET_SLUG}\",               # Datasets tab\n    f\"/kaggle/input/competitions/{DATASET_SLUG}\",  # Competitions tab\n]\n\nMETADATA_CSV = \"train.csv\"\nIMAGE_SUBDIRS = [\"train_images\"]\nID_COLUMN = \"id_code\"\nLABEL_COLUMN = \"diagnosis\"\n\n# APTOS's test half is the competition's scoring holdout: test.csv carries id_code\n# and nothing else — there is no diagnosis column — and test_images is unlabelled.\n# It is excluded outright rather than merely left unused, because both ways of\n# getting this wrong are quiet ones. Point METADATA_CSV at test.csv and the\n# organiser fails deep inside on a column that was never there; add test_images to\n# IMAGE_SUBDIRS and ~3,600 unlabelled images become matchless rows that read like a\n# path bug. The check below refuses either before any of that happens.\n# The folder-scan fallback cannot reach test_images either — it holds no class\n# subdirectories, so there is nothing there for it to match.\nEXCLUDED_CSVS = [\"test.csv\", \"sample_submission.csv\"]\nEXCLUDED_SUBDIRS = [\"test_images\"]\n\n# APTOS images are PNG. Per dataset, never shared: a wrong extension matches zero\n# files, and zero files looks exactly like a wrong path, so the two failures are\n# worth keeping impossible to confuse.\nIMAGE_EXTENSION = \".png\"\n\nCLASS_NAMES = [\"No_DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative_DR\"]\n# Maps raw CSV label values onto CLASS_NAMES. Both datasets grade 0-4 identically.\nLABEL_MAP = {\n    \"0\": \"No_DR\",\n    \"1\": \"Mild\",\n    \"2\": \"Moderate\",\n    \"3\": \"Severe\",\n    \"4\": \"Proliferative_DR\",\n}\n\n# ---------------------------------------------------------------------------\n# VALIDATION DATA — IDRiD  (CLAUDE.md > Datasets: idrid-dataset, validation only)\n# ---------------------------------------------------------------------------\n# A separate cohort, on different equipment, in a different country. Validation\n# ONLY, never merged into training: scoring on held-out APTOS mostly measures how\n# well the model learned APTOS's camera, and a screening tool that works only on\n# the hardware it was trained on is worth discovering before deployment rather\n# than after. If IDRiD is not attached the run falls back to an internal split and\n# records which it used in metrics.json.\n# CONFIRMED against a real Kaggle session: IDRiD mounts under the full-path form,\n# not the bare slug. That path is listed first because it is the one actually\n# observed; the resolver still globs the owner segment behind it, so a teammate\n# who attached a different copy is not left out.\nIDRID_DATASET_SLUG = \"idrid-dataset\"\nIDRID_DATASET_DIR_CANDIDATES = [\n    f\"/kaggle/input/datasets/mariaherrerot/{IDRID_DATASET_SLUG}\",  # confirmed real path\n    f\"/kaggle/input/{IDRID_DATASET_SLUG}\",                         # Datasets tab\n]\nIDRID_METADATA_CSV = \"\"              # \"\" searches for the disease-grading CSV\n# CONFIRMED against the mariaherrerot mirror's idrid_labels.csv: that copy uses\n# APTOS-style headers, not IDRiD's original \"Image name\"/\"Retinopathy grade\".\nIDRID_ID_COLUMN = \"id_code\"\nIDRID_IMAGE_EXTENSION = \".jpg\"       # IDRiD ships JPG where APTOS ships PNG\n\n# IDRiD's grading sheet carries TWO grades per image, for two different diseases:\n# retinopathy and diabetic macular oedema. Naming only the one we want would leave\n# the other as a plausible-looking column a future editor might reach for, so both\n# are named and one is explicitly marked off limits. Training on the DME column\n# would yield a model that scores respectably against the wrong disease.\n#\n# CONFIRMED against the mariaherrerot mirror: its DR grade column is \"diagnosis\".\n# Published copies of IDRiD vary in spacing, capitalisation and naming scheme —\n# the original release calls these \"Image name\" and \"Retinopathy grade\" — so the\n# diagnostic below stays in place: on a different mirror this constant will not\n# match, and the run prints every column in the CSV before falling back.\nIDRID_DR_GRADE_COLUMN = \"diagnosis\"\nIDRID_DME_GRADE_COLUMN = \"Risk of macular edema\"   # NOT a label — different disease\n\n# ---------------------------------------------------------------------------\n# TRAINING\n# ---------------------------------------------------------------------------\nBACKBONE = \"efficientnet_b0\"\nIMAGE_SIZE = 224\nEPOCHS = 6\nBATCH_SIZE = 32\nLEARNING_RATE = 3e-4\n# Kaggle sessions are time-boxed. Raise or set to None for a full run.\nMAX_IMAGES_PER_CLASS = 400\nOUTPUT_DIR = \"/kaggle/working/dr/efficientnet_b0\"\n\n# Ben Graham style crop: trim the black frame around the fundus circle before\n# resizing. A toggle, not a requirement — plain resize trains fine — but on fundus\n# photographs the border is often a third of the frame, so cropping first raises\n# the retina's effective resolution at 224px for essentially no cost. Whatever is\n# set here is written into preprocessing.json, and the backend must match it.\nRETINA_CROP = True\nRETINA_CROP_THRESHOLD = 10\n\n\ndef resolve_dataset_dir() -> pathlib.Path:\n    \"\"\"Return the APTOS root that actually holds METADATA_CSV, else the first candidate.\n\n    Directory existence is not the test — /kaggle/input/competitions can exist\n    because some other competition is attached. The test is whether the labels file\n    is inside it, which is what the rest of the script needs.\n\n    On failure this prints what it looked for and what is actually mounted, in the\n    same spirit as the IDRiD grade-column diagnostic: a wrong path should cost one\n    run and a one-line edit, not a puzzling synthetic-fallback run that looks like\n    the pipeline itself is broken.\n    \"\"\"\n    root = resolve_dataset_root(\n        slug=DATASET_SLUG,\n        marker=METADATA_CSV,\n        extra_candidates=DATASET_DIR_CANDIDATES,\n        label=\"APTOS\",\n    )\n    if root is not None:\n        return root\n    print(\"[warn] add the correct root to DATASET_DIR_CANDIDATES at the top of this script.\")\n    return pathlib.Path(DATASET_DIR_CANDIDATES[0])\n\n\ndef check_excluded_inputs() -> None:\n    \"\"\"Refuse to run if the unlabelled half has been wired in by mistake.\"\"\"\n    if METADATA_CSV in EXCLUDED_CSVS:\n        raise SystemExit(\n            f\"METADATA_CSV is {METADATA_CSV!r}, which is listed in EXCLUDED_CSVS. \"\n            f\"APTOS's test CSV has no {LABEL_COLUMN!r} column — there are no labels \"\n            f\"there to train on. Use train.csv.\")\n    overlap = [d for d in IMAGE_SUBDIRS if d in EXCLUDED_SUBDIRS]\n    if overlap:\n        raise SystemExit(\n            f\"IMAGE_SUBDIRS contains {overlap}, which is listed in EXCLUDED_SUBDIRS. \"\n            f\"Those images are unlabelled; including them cannot help training.\")\n    print(f\"[data] APTOS labelled half only; excluded: {EXCLUDED_CSVS + EXCLUDED_SUBDIRS}\")\n\n\ndef find_idrid_grading_csv(root: pathlib.Path) -> pathlib.Path | None:\n    \"\"\"Locate IDRiD's disease-grading labels among the several CSVs it ships.\n\n    IDRiD bundles grading, segmentation and localisation sheets together, so taking\n    the first CSV found would silently validate against the wrong labels. Prefer a\n    filename mentioning grading, then a training-labels sheet, and announce any\n    fallback — a wrong choice here produces plausible numbers that mean nothing.\n    \"\"\"\n    if IDRID_METADATA_CSV:\n        explicit = root / IDRID_METADATA_CSV\n        if explicit.exists():\n            return explicit\n        print(f\"[warn] IDRID_METADATA_CSV={IDRID_METADATA_CSV!r} not found under {root}\")\n        return None\n    candidates = sorted(root.rglob(\"*.csv\"))\n    if not candidates:\n        return None\n    for keywords in ((\"grading\", \"train\"), (\"grading\",), (\"grade\",), (\"train\",)):\n        for path in candidates:\n            name = path.name.lower()\n            if all(k in name for k in keywords):\n                return path\n    print(f\"[warn] no obvious IDRiD grading CSV; using {candidates[0].name}. \"\n          f\"Set IDRID_METADATA_CSV if that is wrong.\")\n    return candidates[0]\n\n\ndef resolve_idrid_dr_column(csv_path: pathlib.Path) -> str | None:\n    \"\"\"Confirm the DR-grade column exists, or print what the CSV actually has.\"\"\"\n    import pandas as pd\n\n    try:\n        header = list(pd.read_csv(csv_path, nrows=0).columns)\n    except Exception as exc:\n        print(f\"[warn] could not read {csv_path.name}: {type(exc).__name__}: {exc}\")\n        return None\n\n    if IDRID_DR_GRADE_COLUMN in header:\n        return IDRID_DR_GRADE_COLUMN\n\n    print(f\"[warn] {csv_path.name} has no column {IDRID_DR_GRADE_COLUMN!r}.\")\n    print(f\"[warn] available columns: {header}\")\n    likely = [c for c in header\n              if \"retino\" in c.lower()\n              or (\"grade\" in c.lower() and \"edema\" not in c.lower()\n                  and \"oedema\" not in c.lower())]\n    if likely:\n        print(f\"[warn] closest matches for the DR grade: {likely}\")\n    if IDRID_DME_GRADE_COLUMN in header:\n        print(f\"[warn] {IDRID_DME_GRADE_COLUMN!r} IS present, but that is macular \"\n              f\"oedema — a different disease. Do not use it as the DR grade.\")\n    print(\"[warn] set IDRID_DR_GRADE_COLUMN at the top of this script.\")\n    return None\n\n\ndef idrid_class_dirs() -> dict:\n    \"\"\"Build IDRiD's {class: directory} mapping for validation, or {} if absent.\"\"\"\n    root = resolve_dataset_root(\n        slug=IDRID_DATASET_SLUG,\n        extra_candidates=IDRID_DATASET_DIR_CANDIDATES,\n        label=\"IDRiD\",\n    )\n    if root is None:\n        print(\"[data] IDRiD not attached; falling back to an internal split of APTOS.\")\n        return {}\n    csv_path = find_idrid_grading_csv(root)\n    if csv_path is None:\n        print(f\"[warn] no usable CSV under {root}; \"\n              f\"falling back to an internal split of APTOS.\")\n        return {}\n    print(f\"[data] IDRiD validation labels: {csv_path.name}\")\n\n    dr_column = resolve_idrid_dr_column(csv_path)\n    if dr_column is None:\n        print(\"[warn] falling back to an internal split of APTOS.\")\n        return {}\n\n    return organise_from_metadata_csv(\n        csv_path=csv_path,\n        image_dirs=[root],\n        id_column=IDRID_ID_COLUMN,\n        label_column=dr_column,\n        class_names=CLASS_NAMES,\n        destination=working_root() / \"_organised_dr_idrid_val\",\n        label_map=LABEL_MAP,\n        id_suffixes=(IDRID_IMAGE_EXTENSION,),\n    )\n\n\ndef main() -> None:\n    check_excluded_inputs()\n\n    dataset_dir = resolve_dataset_dir()\n\n    config = TrainingConfig(\n        disease_id=\"dr\",\n        dataset_dir=str(dataset_dir),\n        class_names=CLASS_NAMES,\n        backbone=BACKBONE,\n        image_size=IMAGE_SIZE,\n        epochs=EPOCHS,\n        batch_size=BATCH_SIZE,\n        learning_rate=LEARNING_RATE,\n        max_images_per_class=MAX_IMAGES_PER_CLASS,\n        output_dir=OUTPUT_DIR,\n        retina_crop=RETINA_CROP,\n        retina_crop_threshold=RETINA_CROP_THRESHOLD,\n    )\n\n    class_dirs = organise_from_metadata_csv(\n        csv_path=dataset_dir / METADATA_CSV,\n        image_dirs=[dataset_dir / d for d in IMAGE_SUBDIRS],\n        id_column=ID_COLUMN,\n        label_column=LABEL_COLUMN,\n        class_names=CLASS_NAMES,\n        destination=working_root() / \"_organised_dr\",\n        label_map=LABEL_MAP,\n        id_suffixes=(IMAGE_EXTENSION,),\n    )\n\n    train(config, class_dirs=class_dirs or None, val_class_dirs=idrid_class_dirs() or None)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T17:02:56.024106Z","iopub.execute_input":"2026-08-30T17:02:56.024495Z"}},"outputs":[],"execution_count":null}]}