{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"eb3937b2","cell_type":"markdown","source":"\n# OT-WGAN Chest X-ray Augmentation & Classification (4-class, COVID-19 Radiography Database)\n\n**Fourth iteration.** The 5-backbone run completed (all 20 classifier runs finished),\nbut Grad-CAM crashed, and a separate 3-class side-run of this same codebase surfaced\nevidence that the augmentation-group grid was cutting off each model's actual best\nresult. Both fixed here — nothing else changed.\n\n## Fix 1: Grad-CAM crash (`RuntimeError: ...view is being modified inplace`)\n\nVGG19's ReLU layers use `inplace=True` by default, which conflicts with\n`register_full_backward_hook` — a known PyTorch interaction, not a bug specific to\nthis notebook's Grad-CAM code. **Fix**: disable every `inplace`-capable module on the\nmodel right before hooking it (the standard workaround, e.g. used internally by the\n`pytorch-grad-cam` library). This only affects the Grad-CAM forward/backward pass —\nit doesn't touch the model's trained weights or its behavior anywhere else. The\nper-class display loop's exception handling was also widened from catching only\n`KeyError` to catching any exception, so one bad class/layer combination shows an\ninline error label instead of blanking the whole figure.\n\n## Fix 2: augmentation-group grid moved to the range endpoints, not the middle\n\nA separate 3-class run of this same codebase (same GAN, same `train_classifier`,\njust fewer classes and the *original, full* 8-point grid\n`classic=[200,400,700,1000]` / `synthetic=[100,300,600,900]`) surfaced a pattern:\nseveral models' actual best-performing group sat at the **extremes** of the range,\nnot the middle —VGG19 peaked at classic **C=1000** (91.7%), ResNet50 peaked at\nsynthetic **G=900** (92.5%), both well clear of their neighboring points. The\nprevious version's shrunk grid, `classic=[200,700]` / `synthetic=[300,600]`, was an\narbitrary middle subset that would have missed both of those peaks entirely.\n\n**Fix**: test the range endpoints instead — `classic_groups_per_class = [200, 1000]`,\n`synthetic_groups_per_class = [100, 900]`. Same compute cost (still 2 classic + 2\nsynthetic = 4 points per model, 20 runs total), but now anchored on where the\nevidence says the real optima tend to sit, rather than an arbitrary midpoint. Runtime\nimpact is minor and roughly self-canceling: the classic curve's upper point grows\n(700→1000, somewhat more images/epoch) while the synthetic curve's lower point\nshrinks (300→100, fewer images/epoch) and its upper point grows a little (600→900);\nnet effect should keep this in the same 5-6h ballpark as before.\n\n## Otherwise unchanged\n\n## What changed from the 6-backbone run, and why\n\n**1. EfficientNet-B0 dropped.** Its own numbers in the cancelled run were bad —\n**70–76% accuracy**, against 85–92% for every other model — not a \"slightly\nweaker\" backbone, a real underperformance signal. Cutting it saves time *and*\ndrops the worst performer; no tradeoff.\n\n**2. Augmentation groups shrunk — the actual time driver.** Classifier training\n(not GAN training) was ~90% of total runtime, and it scales directly with\n`len(classic_groups_per_class) + len(synthetic_groups_per_class)` per model:\n\n| Setting | Was | Now | Effect |\n|---|---|---|---|\n| `classic_groups_per_class` | `[200,400,700,1000]` | `[200,700]` | 4→2 points on this curve |\n| `synthetic_groups_per_class` | `[100,300,600,900]` | `[300,600]` | 4→2 points on this curve |\n| `classifier_epochs` (ceiling) | 60 | 40 | every run already early-stopped well under 60 |\n| `early_stop_patience` | 7 | 5 | tighter stop, same restored-best-checkpoint logic |\n| `gan_iterations` | 6000 | 4000 | GAN phase is a small fraction of total time, but free savings |\n\nCombos per model: was 8 (4+4), now **4 (2+2)** — combined with dropping one\nbackbone, that's **5 models × 4 groups = 20 training runs**, down from 48. This\nstill keeps one point on each end of both curves (smallest/largest classic,\nsmallest/largest synthetic), so the classic-vs-synthetic comparison the whole\nnotebook is built around still holds — just with 2 points per line instead of 4.\n\n## Everything else is unchanged from what was already working\n\nDataset (`tawsifurrahman/covid19-radiography-database`, 4-class), backbones\n(**VGG19, Inception V3, Xception, ResNet50, DenseNet121**), WGAN-GP, `geomloss`\nSinkhorn divergence on pooled features in both critic and generator loss,\nnearest-neighbor-upsample generator, InstanceNorm critic, 200 real images/class,\nstaged fine-tuning (8-epoch frozen warmup → full unfreeze), weight decay, label\nsmoothing, mixed precision, and the auto-detected Grad-CAM last-conv-layer lookup.\n\n## Still use the resume machinery\n\nEven at ~20 runs this is a multi-hour job. Use Kaggle's **Save & Run All\n(Commit)** so a timeout doesn't lose progress — results are saved to CSV after\nevery single group finishes, and `RECOVERY_MODE=True` next time reloads them and\nretrains only the best group per model for confusion matrices/Grad-CAM, rather\nthan starting over.\n\n## Grad-CAM layer lookup made more robust\n\nThe original notebook hardcoded a `LAST_CONV_LAYER` name per architecture, with a\ncomment flagging that timm's internal naming could change and break it. Since\nDenseNet121 was added and its exact internal layer name wasn't something I could\nverify without running it, Grad-CAM now **auto-detects the last `Conv2d` module**\nby walking `model.named_modules()`, the same robust approach used across all five\nbackbones here — no hardcoded name to get wrong for the new one.\n","metadata":{}},{"id":"badd5866","cell_type":"markdown","source":"## 0. Setup","metadata":{}},{"id":"626397e9","cell_type":"code","source":"\n!pip install -q timm geomloss kaggle\n","metadata":{},"outputs":[],"execution_count":null},{"id":"7525ed0a","cell_type":"code","source":"\nimport os\nimport gc\nimport json\nimport random\nimport shutil\nimport zipfile\nimport datetime\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom IPython.display import display\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nimport torchvision.models as tvm\nfrom torch.utils.data import Dataset, DataLoader\n\nimport timm\nfrom geomloss import SamplesLoss\n\nfrom sklearn.metrics import (\n    accuracy_score, precision_recall_fscore_support,\n    confusion_matrix, roc_curve, auc, classification_report,\n)\nfrom sklearn.model_selection import train_test_split\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"device:\", DEVICE)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b6720180","cell_type":"markdown","source":"## 1. Configuration","metadata":{}},{"id":"5174f94a","cell_type":"code","source":"\n# PAPER_PROTOCOL: values explicitly stated by the ICCIT OT-GAN paper or by the\n# rossettisimone/AUGMENTATION_GAN repo.\nPAPER_PROTOCOL = {\n    \"classes\": [\"COVID\", \"Normal\", \"Lung_Opacity\", \"Viral Pneumonia\"],  # Radiography DB folder names\n    \"gan_image_size\": 64,               # repo's DCGAN resolution\n    \"gan_channels\": 1,                  # repo trains on single-channel (grayscale) images\n    \"gan_latent_dim\": 100,              # repo's nz\n    \"gan_lr_generator\": 0.0002,         # matched G/D lrs for WGAN-GP\n    \"gan_lr_discriminator\": 0.0002,\n    \"critic_weight_clip\": 0.01,         # kept for reference; not used (gradient penalty instead)\n    \"gp_lambda\": 10.0,                  # standard WGAN-GP gradient-penalty weight\n    \"sinkhorn_pool_size\": 16,           # Sinkhorn computed on images pooled to this HxW first\n    \"classic_groups_per_class\": [200, 1000],\n    \"synthetic_groups_per_class\": [100, 900],\n    \"vgg19_input\": 224,                 # kept for reference; VGG19 not in MODEL_NAMES this run\n    \"inceptionv3_input\": 299,\n    \"xception_input\": 299,\n    \"resnet50_input\": 224,              # not in either paper -- ResNet50 is your addition\n    \"efficientnet_b0_input\": 224,       # new backbone, added on request\n    \"densenet121_input\": 224,           # new backbone, added on request -- CheXNet's architecture\n}\n\n# IMPLEMENTATION_ASSUMPTIONS: neither source specifies these for this exact setup, so\n# they're explicit, adjustable defaults rather than silent choices.\nIMPLEMENTATION_ASSUMPTIONS = {\n    \"real_train_per_class\": 200,    # bumped from 150 -- the smallest class in this dataset\n                                     # (Viral Pneumonia, ~1345 total) comfortably supports more\n                                     # than prashant268's smallest class did; check cell 8's\n                                     # printed per-class counts before raising further\n    \"real_test_per_class\": 50,       # bumped from 40 for the same reason\n    \"gan_iterations\": 4000,\n    \"gan_batch_size\": 16,\n    \"n_critic\": 5,                  # standard WGAN: critic steps per generator step\n    \"sinkhorn_blur\": 0.05,          # geomloss entropic regularization strength (epsilon)\n    \"sinkhorn_lambda_critic\": 0.1,  # weight of Sinkhorn term added to the critic loss\n    \"sinkhorn_lambda_generator\": 0.1,  # weight of Sinkhorn term added to the generator loss\n    \"synthetic_pool_per_class\": 2000,  # generate enough images to cover the largest group\n    \"classifier_batch_size\": 32,\n    \"classifier_epochs\": 40,        # this is a CEILING, not a target -- early stopping\n                                     # (see below) will stop most runs well before this\n    \"early_stop_patience\": 5,\n    \"early_stop_min_delta\": 0.003,\n    \"use_amp\": True,                # mixed-precision training -- ~1.5-2x GPU speedup, no accuracy cost\n    \"classifier_lr\": 1e-4,\n    \"classifier_weight_decay\": 1e-4,\n    \"label_smoothing\": 0.05,\n    \"val_frac\": 0.15,\n    \"use_staged_finetuning\": True,  # freeze backbone for a head-only warmup, then unfreeze\n    \"finetune_warmup_epochs\": 8,\n    \"finetune_lr\": 2e-5,\n    \"quick_debug\": False,           # FULL RUN: real GAN iterations, real epoch counts, real group sizes\n}\n\nCLASSES = PAPER_PROTOCOL[\"classes\"]\nNUM_CLASSES = len(CLASSES)\nCLASS_TO_IDX = {c: i for i, c in enumerate(CLASSES)}\n\nif IMPLEMENTATION_ASSUMPTIONS[\"quick_debug\"]:\n    QUICK = {\n        \"real_train_per_class\": 20,\n        \"real_test_per_class\": 10,\n        \"gan_iterations\": 150,\n        \"synthetic_pool_per_class\": 60,\n        \"classic_groups_per_class\": [20, 40],\n        \"synthetic_groups_per_class\": [10, 30],\n        \"classifier_epochs\": 2,\n        \"finetune_warmup_epochs\": 1,\n    }\n    IMPLEMENTATION_ASSUMPTIONS.update({k: v for k, v in QUICK.items()\n                                        if k in IMPLEMENTATION_ASSUMPTIONS})\n    PAPER_PROTOCOL[\"classic_groups_per_class\"] = QUICK[\"classic_groups_per_class\"]\n    PAPER_PROTOCOL[\"synthetic_groups_per_class\"] = QUICK[\"synthetic_groups_per_class\"]\n    print(\"QUICK_DEBUG mode is ON — tiny epoch/sample counts. \"\n          \"Set IMPLEMENTATION_ASSUMPTIONS['quick_debug']=False for a full run.\")\n\nWORK_DIR = Path(\"./otgan_work\")\nWORK_DIR.mkdir(exist_ok=True)\n\n# --- RECOVERY MODE ---\n# Set True to reuse a previous run's saved files (skips GAN training and the full\n# 40-combo grid; only retrains the single best group per model for confusion\n# matrices/Grad-CAM). Set False for a genuine first-time full run from scratch.\n# Defaulting to False here: this dataset+model combination has no prior saved run to\n# recover from yet. Once a full run completes and its otgan_work/ folder is saved as\n# a Kaggle output/input, flip this to True and point RECOVERY_SOURCE_DIR at it.\nRECOVERY_MODE = False\nRECOVERY_SOURCE_DIR = Path(\"/kaggle/input/PASTE_YOUR_SAVED_RUN_PATH_HERE/otgan_work\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"90a2b0f7","cell_type":"code","source":"# Copy a previous run's saved files into WORK_DIR (only runs if RECOVERY_MODE is on).\n# If your attached input dataset's path differs from RECOVERY_SOURCE_DIR above, edit\n# it there (check the \"Input\" panel on the right for the exact mounted path).\nif RECOVERY_MODE:\n    assert RECOVERY_SOURCE_DIR.exists(), (\n        f\"RECOVERY_SOURCE_DIR not found: {RECOVERY_SOURCE_DIR} -- check the Input panel \"\n        \"for the real mounted path, or set RECOVERY_MODE=False for a fresh run.\"\n    )\n    copied = []\n    for f in RECOVERY_SOURCE_DIR.glob(\"*\"):\n        shutil.copy(f, WORK_DIR / f.name)\n        copied.append(f.name)\n    print(f\"RECOVERY_MODE on -- copied {len(copied)} files into {WORK_DIR}:\")\n    for name in sorted(copied):\n        print(\" -\", name)\nelse:\n    print(\"RECOVERY_MODE off -- starting a fresh run, nothing copied.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"ca989264","cell_type":"markdown","source":"\n## 2. Download the dataset\n\nNow using `tawsifurrahman/covid19-radiography-database` (4 classes: `COVID`,\n`Normal`, `Lung_Opacity`, `Viral Pneumonia`, each in its own folder with an\n`images/` subfolder — no train/test split provided, so this notebook builds its\nown balanced split exactly like it did for `prashant268`). Upload your\n`kaggle.json` API token first (Colab), or attach the dataset directly (Kaggle\nnotebooks) and skip the download cell.\n","metadata":{}},{"id":"6bc1e0a8","cell_type":"code","source":"\nKAGGLE_DATASET = \"tawsifurrahman/covid19-radiography-database\"\nDATA_ROOT = Path(\"/kaggle/input/datasets/tawsifurrahman/covid19-radiography-database\")\n\n# The Radiography DB's own folder names occasionally vary by version (space vs\n# underscore, casing) -- resolve fuzzily rather than hardcoding one exact spelling.\nCLASS_FOLDER_VARIANTS = {\n    \"COVID\": [\"COVID\", \"COVID-19\", \"COVID19\"],\n    \"Normal\": [\"Normal\"],\n    \"Lung_Opacity\": [\"Lung_Opacity\", \"Lung Opacity\"],\n    \"Viral Pneumonia\": [\"Viral Pneumonia\", \"Viral_Pneumonia\"],\n}\n\ndef download_kaggle_dataset():\n    os.environ.setdefault(\"KAGGLE_CONFIG_DIR\", str(Path(\"~/.kaggle\").expanduser()))\n    if not DATA_ROOT.exists():\n        DATA_ROOT.mkdir(parents=True)\n        os.system(f\"kaggle datasets download -d {KAGGLE_DATASET} -p {DATA_ROOT} --unzip\")\n    return DATA_ROOT\n\n# Uncomment when kaggle.json is configured in this environment:\n# download_kaggle_dataset()\n\ndef resolve_class_dir(root, variants):\n    '''Find this class's folder under `root` (case/space/underscore-insensitive),\n    preferring an `images/` subfolder if present (this dataset nests images one\n    level deeper than the class folder, unlike prashant268).'''\n    if not root.is_dir():\n        return None\n    for entry in os.listdir(root):\n        entry_norm = entry.lower().replace(\"_\", \" \").strip()\n        for variant in variants:\n            if entry_norm == variant.lower().replace(\"_\", \" \").strip():\n                full = root / entry\n                images_sub = full / \"images\"\n                return images_sub if images_sub.is_dir() else full\n    return None\n\ndef locate_dataset_root(start=DATA_ROOT):\n    '''Find the folder that contains all 4 class subfolders.'''\n    for root, dirs, _files in os.walk(start):\n        root_p = Path(root)\n        found = {cls: resolve_class_dir(root_p, variants)\n                 for cls, variants in CLASS_FOLDER_VARIANTS.items()}\n        if all(v is not None for v in found.values()):\n            return root_p\n    return None\n\nDATASET_ROOT = locate_dataset_root()\nprint(\"Located dataset root:\", DATASET_ROOT)\n\nCLASS_DIRS = {}\nif DATASET_ROOT is not None:\n    for cls, variants in CLASS_FOLDER_VARIANTS.items():\n        d = resolve_class_dir(DATASET_ROOT, variants)\n        CLASS_DIRS[cls] = d\n        n = len(list(d.glob(\"*\"))) if d else 0\n        print(f\"  {cls}: {d}  ({n} images)\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"05dd8c90","cell_type":"markdown","source":"\n## 3. Build the fixed, disjoint real train/test split\n\nFollowing the same protocol as the 3-class run: a small **balanced** real\ntraining pool is sampled per class, and a separate, **disjoint** real test pool\nis held out. The test pool is never touched by the GAN, by classic augmentation,\nor by classifier-selection decisions — only used for final evaluation. Checked\nin the leakage audit at the end.\n","metadata":{}},{"id":"320a8b2e","cell_type":"code","source":"\ndef list_all_images(dataset_root, classes):\n    '''Pool every image file per class directly from its resolved folder -- unlike\n    prashant268, this dataset has no train/test split to pool across; CLASS_DIRS\n    (cell 8) already points straight at each class's image folder.'''\n    pools = {}\n    for cls in classes:\n        d = CLASS_DIRS.get(cls)\n        pools[cls] = sorted(d.glob(\"*\")) if d is not None else []\n    return pools\n\n\ndef build_balanced_split(dataset_root, classes, n_train, n_test, seed=SEED):\n    rng = random.Random(seed)\n    pools = list_all_images(dataset_root, classes)\n    split = {\"train\": {}, \"test\": {}}\n    for cls in classes:\n        files = pools[cls][:]\n        rng.shuffle(files)\n        need = n_train + n_test\n        assert len(files) >= need, (\n            f\"class {cls} has only {len(files)} images, need {need}. \"\n            f\"Lower real_train_per_class/real_test_per_class.\"\n        )\n        split[\"train\"][cls] = files[:n_train]\n        split[\"test\"][cls] = files[n_train:need]\n    return split\n\n\ndef load_gray64(path, size=64):\n    img = Image.open(path).convert(\"L\").resize((size, size), Image.BILINEAR)\n    return np.array(img, dtype=np.uint8)\n\n\nif DATASET_ROOT is not None:\n    SPLIT = build_balanced_split(\n        DATASET_ROOT, CLASSES,\n        n_train=IMPLEMENTATION_ASSUMPTIONS[\"real_train_per_class\"],\n        n_test=IMPLEMENTATION_ASSUMPTIONS[\"real_test_per_class\"],\n    )\n\n    manifest_rows = []\n    for split_name, per_class in SPLIT.items():\n        for cls, files in per_class.items():\n            for f in files:\n                manifest_rows.append({\"split\": split_name, \"class\": cls, \"path\": str(f)})\n    split_manifest = pd.DataFrame(manifest_rows)\n    split_manifest.to_csv(WORK_DIR / \"split_manifest.csv\", index=False)\n\n    train_paths = set(split_manifest.loc[split_manifest[\"split\"] == \"train\", \"path\"])\n    test_paths = set(split_manifest.loc[split_manifest[\"split\"] == \"test\", \"path\"])\n    assert train_paths.isdisjoint(test_paths), \"train/test overlap detected!\"\n    print(\"Balanced split built. Train/test are disjoint.\")\n    print(split_manifest.groupby([\"split\", \"class\"]).size())\nelse:\n    print(\"Dataset not found yet -- run the download cell above first.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"cf6078b7","cell_type":"markdown","source":"\n## 4. GAN training images (grayscale 64x64, real training images only)\n\n**Revised.** The first version used a `DataLoader` with `drop_last=True` and\n`n_critic=5`: with only ~65 real images/class and `batch_size=16`, that's just 4\nbatches per epoch, so the critic-update loop hit `StopIteration` after 4 batches and\nthe generator was updated **once per epoch** — 200 \"epochs\" meant only 200 total\ngenerator updates. WGAN-GP typically needs 10,000+ generator iterations even on\nsmall datasets, so this was silently starving the generator regardless of loss\nfunction or architecture, and is very likely the main reason images looked like\ntextured noise with no anatomical structure.\n\nTraining length is now defined directly as a number of **iterations**\n(`gan_iterations`), each iteration sampling a fresh random batch **with\nreplacement** straight from the small real-image pool (held in memory — 65 images\nat 64x64 is trivial). This is standard practice for GAN training on small datasets\nand fully decouples training length from how few real images there are.\n","metadata":{}},{"id":"2ab763fa","cell_type":"code","source":"\ndef load_class_tensor(class_name, size=64):\n    '''Load every real TRAINING image for one class into memory as a single tensor,\n    normalized to [-1, 1] to match the generator's Tanh output.'''\n    paths = SPLIT[\"train\"][class_name]\n    arr = np.stack([load_gray64(p, size=size) for p in paths]).astype(np.float32)\n    arr = (arr / 127.5) - 1.0\n    return torch.from_numpy(arr).unsqueeze(1)  # (N, 1, H, W)\n\n\ndef sample_real_batch(real_tensor, batch_size, device):\n    '''Random batch WITH replacement -- decouples batch composition from the size\n    of the (tiny) real image pool.'''\n    idx = torch.randint(0, real_tensor.size(0), (batch_size,))\n    return real_tensor[idx].to(device)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c2f9de03","cell_type":"markdown","source":"\n## 5. OT-WGAN architecture\n\n**Revised from the first version of this notebook** after inspecting its output:\ngenerated images showed a strong checkerboard/waffle grid and no real anatomical\nstructure. Two architecture-level causes, independent of training length:\n\n- `ConvTranspose2d` upsampling (the repo's original choice) produces well-documented\n  checkerboard artifacts from uneven kernel/stride overlap. The generator below uses\n  **nearest-neighbor upsample + Conv2d** instead, which doesn't have this failure mode.\n- `BatchNorm2d` inside the critic correlates statistics *within* a batch, which\n  conflicts with the per-sample gradient penalty used below (WGAN-GP). The critic now\n  uses `InstanceNorm2d` instead, which is the standard recommendation for WGAN-GP.\n","metadata":{}},{"id":"a496ac99","cell_type":"code","source":"\nNZ = PAPER_PROTOCOL[\"gan_latent_dim\"]\nNGF = 64\nNDF = 64\nNC = PAPER_PROTOCOL[\"gan_channels\"]\n\n\ndef weights_init(m):\n    classname = m.__class__.__name__\n    if classname.find(\"Conv\") != -1:\n        nn.init.normal_(m.weight.data, 0.0, 0.02)\n    elif classname.find(\"BatchNorm\") != -1 or classname.find(\"InstanceNorm\") != -1:\n        if getattr(m, \"weight\", None) is not None:\n            nn.init.normal_(m.weight.data, 1.0, 0.02)\n        if getattr(m, \"bias\", None) is not None:\n            nn.init.constant_(m.bias.data, 0)\n\n\ndef upsample_block(in_ch, out_ch):\n    '''Nearest-neighbor upsample + Conv2d: avoids the checkerboard artifact that\n    ConvTranspose2d(kernel=4, stride=2) produces.'''\n    return nn.Sequential(\n        nn.Upsample(scale_factor=2, mode=\"nearest\"),\n        nn.Conv2d(in_ch, out_ch, 3, 1, 1, bias=False),\n        nn.BatchNorm2d(out_ch),\n        nn.ReLU(True),\n    )\n\n\nclass Generator(nn.Module):\n    '''nz -> 64x64x1 via a dense projection to 4x4 followed by upsample+conv blocks\n    (revised from the repo's ConvTranspose2d stack -- see markdown above).'''\n    def __init__(self):\n        super().__init__()\n        self.project = nn.Sequential(\n            nn.ConvTranspose2d(NZ, NGF * 8, 4, 1, 0, bias=False),  # 1x1 -> 4x4 only\n            nn.BatchNorm2d(NGF * 8), nn.ReLU(True),\n        )\n        self.upsample_blocks = nn.Sequential(\n            upsample_block(NGF * 8, NGF * 4),   # 4x4   -> 8x8\n            upsample_block(NGF * 4, NGF * 2),   # 8x8   -> 16x16\n            upsample_block(NGF * 2, NGF),       # 16x16 -> 32x32\n        )\n        self.to_image = nn.Sequential(\n            nn.Upsample(scale_factor=2, mode=\"nearest\"),  # 32x32 -> 64x64\n            nn.Conv2d(NGF, NC, 3, 1, 1, bias=False),\n            nn.Tanh(),\n        )\n\n    def forward(self, z):\n        x = self.project(z)\n        x = self.upsample_blocks(x)\n        return self.to_image(x)\n\n\nclass Critic(nn.Module):\n    '''Repo's discriminator conv stack, WGAN-GP-ified: InstanceNorm instead of\n    BatchNorm (BatchNorm is unsound with a per-sample gradient penalty), linear\n    scalar output instead of Sigmoid, no weight clipping (see gradient_penalty below).'''\n    def __init__(self):\n        super().__init__()\n        self.main = nn.Sequential(\n            nn.Conv2d(NC, NDF, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Conv2d(NDF, NDF * 2, 4, 2, 1, bias=False),\n            nn.InstanceNorm2d(NDF * 2, affine=True), nn.LeakyReLU(0.2, inplace=True),\n            nn.Conv2d(NDF * 2, NDF * 4, 4, 2, 1, bias=False),\n            nn.InstanceNorm2d(NDF * 4, affine=True), nn.LeakyReLU(0.2, inplace=True),\n            nn.Conv2d(NDF * 4, NDF * 8, 4, 2, 1, bias=False),\n            nn.InstanceNorm2d(NDF * 8, affine=True), nn.LeakyReLU(0.2, inplace=True),\n            nn.Conv2d(NDF * 8, 1, 4, 1, 0, bias=False),\n        )\n\n    def forward(self, x):\n        return self.main(x).view(-1)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"2d150f8d","cell_type":"markdown","source":"\n## 6. Wasserstein critic loss (with gradient penalty) + pooled Sinkhorn divergence\n\n**Revised from the first version.** Two changes after inspecting the first run's output:\n\n- **Weight clipping -> gradient penalty (WGAN-GP).** Clipping critic weights to\n  `[-0.01, 0.01]` (the paper's stated clamp range) is the original, weak WGAN\n  stabilization trick and is well known to under-power the critic, capping how much\n  signal it can ever give the generator. It's replaced here with the standard WGAN-GP\n  gradient penalty, which enforces the same 1-Lipschitz constraint without that\n  capacity loss.\n- **Sinkhorn on pooled features, not raw pixels.** Computing Sinkhorn divergence on\n  raw flattened 64x64 pixels (4096-dim Euclidean space) is a poor perceptual metric —\n  it mostly captures gross intensity/contrast statistics, not structure. Images are\n  now average-pooled to `sinkhorn_pool_size x sinkhorn_pool_size` first (matching the\n  approach in your leakage-free reference notebook) before the divergence is computed.\n\nThe paper doesn't specify how Wasserstein and Sinkhorn combine, so Sinkhorn is still\nadded as a regularization term on top of the Wasserstein critic loss, per the paper's\nown description of its stabilizing role.\n","metadata":{}},{"id":"191cd51e","cell_type":"code","source":"\nsinkhorn_fn = SamplesLoss(\n    loss=\"sinkhorn\", p=2,\n    blur=IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_blur\"],\n    debias=True,\n)\n\nPOOL_SIZE = PAPER_PROTOCOL[\"sinkhorn_pool_size\"]\n\ndef sinkhorn_divergence(real_batch, fake_batch):\n    '''Average-pool to POOL_SIZE x POOL_SIZE before computing the divergence --\n    raw-pixel Euclidean distance is a poor perceptual metric and made this term\n    nearly meaningless in the first version of this notebook.'''\n    real_pooled = F.adaptive_avg_pool2d(real_batch, POOL_SIZE)\n    fake_pooled = F.adaptive_avg_pool2d(fake_batch, POOL_SIZE)\n    real_flat = real_pooled.reshape(real_pooled.size(0), -1)\n    fake_flat = fake_pooled.reshape(fake_pooled.size(0), -1)\n    return sinkhorn_fn(real_flat, fake_flat)\n\n\ndef gradient_penalty(critic, real_batch, fake_batch):\n    '''Standard WGAN-GP gradient penalty: enforces the 1-Lipschitz constraint on the\n    critic by penalizing gradient norms that deviate from 1 along lines interpolated\n    between real and fake samples. Replaces the paper's weight clipping, which is\n    known to under-power the critic.'''\n    b_size = real_batch.size(0)\n    eps = torch.rand(b_size, 1, 1, 1, device=real_batch.device)\n    interpolated = (eps * real_batch + (1 - eps) * fake_batch).requires_grad_(True)\n    scores = critic(interpolated)\n    grads = torch.autograd.grad(\n        outputs=scores, inputs=interpolated,\n        grad_outputs=torch.ones_like(scores),\n        create_graph=True,\n    )[0]\n    grads = grads.view(b_size, -1)\n    penalty = ((grads.norm(2, dim=1) - 1.0) ** 2).mean()\n    return penalty\n\n\ndef critic_step(generator, critic, opt_critic, real_batch, lambda_sinkhorn, lambda_gp):\n    b_size = real_batch.size(0)\n    z = torch.randn(b_size, NZ, 1, 1, device=DEVICE)\n    with torch.no_grad():\n        fake_batch = generator(z)\n\n    opt_critic.zero_grad()\n    wasserstein_term = critic(fake_batch).mean() - critic(real_batch).mean()\n    gp_term = gradient_penalty(critic, real_batch, fake_batch)\n    sinkhorn_term = sinkhorn_divergence(real_batch, fake_batch)\n    loss_critic = wasserstein_term + lambda_gp * gp_term + lambda_sinkhorn * sinkhorn_term\n    loss_critic.backward()\n    opt_critic.step()\n\n    with torch.no_grad():\n        d_x = critic(real_batch).mean().item()\n        d_g_z1 = critic(fake_batch).mean().item()\n    return loss_critic.item(), d_x, d_g_z1\n\n\ndef generator_step(generator, critic, opt_generator, real_batch, lambda_sinkhorn):\n    b_size = real_batch.size(0)\n    z = torch.randn(b_size, NZ, 1, 1, device=DEVICE)\n    fake_batch = generator(z)\n\n    opt_generator.zero_grad()\n    wasserstein_term = -critic(fake_batch).mean()\n    sinkhorn_term = sinkhorn_divergence(real_batch, fake_batch)\n    loss_generator = wasserstein_term + lambda_sinkhorn * sinkhorn_term\n    loss_generator.backward()\n    opt_generator.step()\n\n    with torch.no_grad():\n        d_g_z2 = critic(fake_batch).mean().item()\n    return loss_generator.item(), d_g_z2, sinkhorn_term.item()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"aee00dc5","cell_type":"markdown","source":"## 7. Train one OT-WGAN per class, then generate the synthetic image pool","metadata":{}},{"id":"3771b006","cell_type":"code","source":"\ndef train_ot_wgan(class_name, iterations, batch_size, log_every=None):\n    real_tensor = load_class_tensor(class_name, size=PAPER_PROTOCOL[\"gan_image_size\"])\n    generator = Generator().to(DEVICE)\n    critic = Critic().to(DEVICE)\n    generator.apply(weights_init)\n    critic.apply(weights_init)\n\n    opt_g = torch.optim.Adam(generator.parameters(),\n                              lr=PAPER_PROTOCOL[\"gan_lr_generator\"], betas=(0.5, 0.9))\n    opt_c = torch.optim.Adam(critic.parameters(),\n                              lr=PAPER_PROTOCOL[\"gan_lr_discriminator\"], betas=(0.5, 0.9))\n\n    n_critic = IMPLEMENTATION_ASSUMPTIONS[\"n_critic\"]\n    lam_c = IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_lambda_critic\"]\n    lam_g = IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_lambda_generator\"]\n    lam_gp = PAPER_PROTOCOL[\"gp_lambda\"]\n\n    history = {\"iteration\": [], \"loss_critic\": [], \"loss_generator\": [],\n               \"D_x\": [], \"D_G_z1\": [], \"D_G_z2\": [], \"sinkhorn\": []}\n\n    for it in range(iterations):\n        for _ in range(n_critic):\n            real_batch = sample_real_batch(real_tensor, batch_size, DEVICE)\n            lc, d_x, d_g_z1 = critic_step(generator, critic, opt_c, real_batch, lam_c, lam_gp)\n\n        real_batch = sample_real_batch(real_tensor, batch_size, DEVICE)\n        lg, d_g_z2, sink = generator_step(generator, critic, opt_g, real_batch, lam_g)\n\n        history[\"iteration\"].append(it)\n        history[\"loss_critic\"].append(lc)\n        history[\"loss_generator\"].append(lg)\n        history[\"D_x\"].append(d_x)\n        history[\"D_G_z1\"].append(d_g_z1)\n        history[\"D_G_z2\"].append(d_g_z2)\n        history[\"sinkhorn\"].append(sink)\n\n        if log_every and (it + 1) % log_every == 0:\n            print(f\"[{class_name}] iter {it+1}/{iterations} \"\n                  f\"loss_C={lc:.3f} loss_G={lg:.3f} \"\n                  f\"D_x={d_x:.3f} D_G_z1={d_g_z1:.3f} D_G_z2={d_g_z2:.3f} \"\n                  f\"sinkhorn={sink:.3f}\")\n\n    return generator, critic, history\n\n\n@torch.no_grad()\ndef generate_synthetic_pool(generator, n_images, batch_size=64):\n    generator.eval()\n    images = []\n    for start in range(0, n_images, batch_size):\n        cur = min(batch_size, n_images - start)\n        z = torch.randn(cur, NZ, 1, 1, device=DEVICE)\n        fake = generator(z).cpu().numpy()\n        fake = ((fake + 1.0) * 127.5).clip(0, 255).astype(np.uint8)  # back to [0,255]\n        images.append(fake)\n    generator.train()\n    return np.concatenate(images, axis=0)  # (N, 1, 64, 64)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"bef5b02e","cell_type":"code","source":"\ntrained_gans = {}\nsynthetic_pools = {}\n\ndef _load_synth_pool(cls):\n    for candidate in (cls, cls[:6]):  # handles filename truncation (e.g. PNEUMONIA -> PNEUMO)\n        p = WORK_DIR / f\"synthetic_pool_{candidate}.npy\"\n        if p.exists():\n            return np.load(p)\n    return None\n\nif RECOVERY_MODE:\n    print(\"RECOVERY_MODE on -- loading synthetic pools from disk instead of training GANs.\")\n    for cls in CLASSES:\n        pool = _load_synth_pool(cls)\n        if pool is None:\n            raise FileNotFoundError(\n                f\"RECOVERY_MODE is on but no saved synthetic_pool_*.npy found for class {cls} \"\n                f\"in {WORK_DIR}. Check that the copy cell above actually found your files, or \"\n                f\"set RECOVERY_MODE=False to train from scratch.\"\n            )\n        synthetic_pools[cls] = pool\n        print(f\"  loaded synthetic pool for {cls}: shape {pool.shape}\")\nelif DATASET_ROOT is not None:\n    for cls in CLASSES:\n        print(f\"=== Training OT-WGAN for class: {cls} ===\")\n        gen, crit, hist = train_ot_wgan(\n            cls,\n            iterations=IMPLEMENTATION_ASSUMPTIONS[\"gan_iterations\"],\n            batch_size=IMPLEMENTATION_ASSUMPTIONS[\"gan_batch_size\"],\n            log_every=max(1, IMPLEMENTATION_ASSUMPTIONS[\"gan_iterations\"] // 10),\n        )\n        trained_gans[cls] = {\"generator\": gen, \"critic\": crit, \"history\": hist}\n\n        pool = generate_synthetic_pool(\n            gen, IMPLEMENTATION_ASSUMPTIONS[\"synthetic_pool_per_class\"]\n        )\n        synthetic_pools[cls] = pool\n        np.save(WORK_DIR / f\"synthetic_pool_{cls}.npy\", pool)\n\n        torch.save({\n            \"generator_state_dict\": gen.state_dict(),\n            \"critic_state_dict\": crit.state_dict(),\n            \"history\": hist,\n        }, WORK_DIR / f\"ot_wgan_{cls}.pt\")\n\n        plt.figure(figsize=(5, 3))\n        plt.plot(hist[\"loss_critic\"], label=\"critic loss\")\n        plt.plot(hist[\"loss_generator\"], label=\"generator loss\")\n        plt.plot(hist[\"sinkhorn\"], label=\"sinkhorn divergence\")\n        plt.legend(); plt.title(f\"OT-WGAN training curves — {cls}\")\n        plt.xlabel(\"iteration\"); plt.tight_layout(); plt.show(); plt.close()\nelse:\n    print(\"Dataset not found yet -- run the download cell above first.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"79e37d6d","cell_type":"markdown","source":"## 8. Real vs. synthetic examples per class","metadata":{}},{"id":"dc026229","cell_type":"code","source":"\nif DATASET_ROOT is not None:\n    fig, axes = plt.subplots(len(CLASSES), 6, figsize=(13, 2.2 * len(CLASSES)))\n    for row, cls in enumerate(CLASSES):\n        real_path = SPLIT[\"train\"][cls][0]\n        real_img = load_gray64(real_path)\n        axes[row, 0].imshow(real_img, cmap=\"gray\")\n        axes[row, 0].set_title(\"real\")\n        axes[row, 0].axis(\"off\")\n        axes[row, 0].set_ylabel(cls, rotation=0, labelpad=40)\n        for col in range(1, 6):\n            axes[row, col].imshow(synthetic_pools[cls][col - 1, 0], cmap=\"gray\")\n            axes[row, col].set_title(\"synthetic\")\n            axes[row, col].axis(\"off\")\n    plt.suptitle(\"Real training examples vs. OT-WGAN generated images\")\n    plt.tight_layout()\n    plt.show(); plt.close()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"510d6af1","cell_type":"markdown","source":"\n## 9. Classic data augmentation (repo's transform pipeline)\n\nMatches `DATASET_BUILDER.ipynb`'s transform list: color jitter, random rotation,\nrandom horizontal flip, random affine (translate + scale), applied to the small\nreal **training** images only.\n","metadata":{}},{"id":"3d31fee0","cell_type":"code","source":"\nclassic_transform = T.Compose([\n    T.RandomApply([T.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.5)], p=0.5),\n    T.RandomApply([T.RandomRotation(degrees=(-5, 5))], p=0.5),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomApply([T.RandomAffine(degrees=0, scale=(0.9, 1.1))], p=0.5),\n    T.RandomApply([T.RandomAffine(degrees=0, translate=(0.1, 0.1))], p=0.5),\n])\n\n\ndef build_classic_augmented_set(real_paths, target_count, size=64, seed=SEED):\n    '''Repeatedly re-applies the classic transform pipeline to the small real pool\n    until target_count augmented images are produced (matches the repo's approach of\n    oversampling augmentations from a small real set).'''\n    rng = random.Random(seed)\n    out = np.empty((target_count, size, size), dtype=np.uint8)\n    for i in range(target_count):\n        path = rng.choice(real_paths)\n        img = Image.open(path).convert(\"L\").resize((size, size), Image.BILINEAR)\n        img = classic_transform(img)\n        out[i] = np.array(img, dtype=np.uint8)\n    return out\n","metadata":{},"outputs":[],"execution_count":null},{"id":"7579873f","cell_type":"markdown","source":"\n## 10. Assemble augmentation groups\n\nRepo's naming: `O1` = the small original real trainset, `C1..C4` = classic-augmented\ngroups of increasing size, `G1..G4` = synthetic-augmented groups of increasing size.\nTwo curves are built for comparison (matching the repo's Fig. 6.13):\n\n- **Classic curve**: `O1 + Ci` for each classic group size (no synthetic images).\n- **Synthetic curve**: `O1 + C_best + Gj` for each synthetic group size, where\n  `C_best` is the largest classic group (the repo added synthetic images on top of\n  the best classic-augmented set, not instead of it).\n","metadata":{}},{"id":"46d10cce","cell_type":"code","source":"\ndef sample_from_pool(pool, n, seed=SEED):\n    rng = np.random.default_rng(seed)\n    idx = rng.choice(len(pool), size=n, replace=(n > len(pool)))\n    return pool[idx, 0]  # drop channel dim -> (n, H, W)\n\n\ndef assemble_dataset(class_image_dict):\n    '''class_image_dict: {class_name: np.ndarray of shape (N, H, W)} -> (X, y) arrays.'''\n    X_parts, y_parts = [], []\n    for cls, imgs in class_image_dict.items():\n        X_parts.append(imgs)\n        y_parts.append(np.full(len(imgs), CLASS_TO_IDX[cls], dtype=np.int64))\n    X = np.concatenate(X_parts, axis=0)\n    y = np.concatenate(y_parts, axis=0)\n    return X, y\n\n\ndef real_only_images():\n    size = PAPER_PROTOCOL[\"gan_image_size\"]\n    return {cls: np.stack([load_gray64(p, size) for p in SPLIT[\"train\"][cls]])\n            for cls in CLASSES}\n\n\ndef real_test_images():\n    size = PAPER_PROTOCOL[\"gan_image_size\"]\n    return {cls: np.stack([load_gray64(p, size) for p in SPLIT[\"test\"][cls]])\n            for cls in CLASSES}\n\n\nif DATASET_ROOT is not None:\n    REAL_TRAIN_IMAGES = real_only_images()\n    REAL_TEST_IMAGES = real_test_images()\n\n    classic_groups = PAPER_PROTOCOL[\"classic_groups_per_class\"]\n    synthetic_groups = PAPER_PROTOCOL[\"synthetic_groups_per_class\"]\n\n    # Classic curve: O1 + Ci (real + increasing classic augmentation, no synthetic)\n    classic_curve_datasets = {}\n    for n in classic_groups:\n        combo = {}\n        for cls in CLASSES:\n            classic_imgs = build_classic_augmented_set(SPLIT[\"train\"][cls], n)\n            combo[cls] = np.concatenate([REAL_TRAIN_IMAGES[cls], classic_imgs], axis=0)\n        classic_curve_datasets[n] = assemble_dataset(combo)\n        print(f\"classic group C={n}: total images = {len(classic_curve_datasets[n][1])}\")\n\n    # Best classic group = largest one, reused as the base for the synthetic curve\n    best_classic_n = max(classic_groups)\n    best_classic_imgs = {\n        cls: build_classic_augmented_set(SPLIT[\"train\"][cls], best_classic_n)\n        for cls in CLASSES\n    }\n\n    # Synthetic curve: O1 + C_best + Gj (increasing synthetic augmentation on top)\n    synthetic_curve_datasets = {}\n    for n in synthetic_groups:\n        combo = {}\n        for cls in CLASSES:\n            synth_imgs = sample_from_pool(synthetic_pools[cls], n)\n            combo[cls] = np.concatenate(\n                [REAL_TRAIN_IMAGES[cls], best_classic_imgs[cls], synth_imgs], axis=0\n            )\n        synthetic_curve_datasets[n] = assemble_dataset(combo)\n        print(f\"synthetic group G={n} (on top of C={best_classic_n}): \"\n              f\"total images = {len(synthetic_curve_datasets[n][1])}\")\n\n    X_test, y_test = assemble_dataset(REAL_TEST_IMAGES)\n    print(\"held-out real test set:\", X_test.shape, y_test.shape)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"673f8a12","cell_type":"markdown","source":"\n## 11. Classifier backbones (transfer learning)\n\nVGG19, InceptionV3, Xception (`timm`), and ResNet50 — all ImageNet-pretrained,\nfine-tuned end-to-end. Grayscale 64x64 arrays are upsampled to each model's native\ninput size and replicated to 3 channels, then ImageNet-normalized.\n","metadata":{}},{"id":"52c0d631","cell_type":"code","source":"\nMODEL_INPUT_SIZES = {\n    \"vgg19\": PAPER_PROTOCOL[\"vgg19_input\"],\n    \"resnet50\": PAPER_PROTOCOL[\"resnet50_input\"],\n    \"inception_v3\": PAPER_PROTOCOL[\"inceptionv3_input\"],\n    \"xception\": PAPER_PROTOCOL[\"xception_input\"],\n    \"efficientnet_b0\": PAPER_PROTOCOL[\"efficientnet_b0_input\"],\n    \"densenet121\": PAPER_PROTOCOL[\"densenet121_input\"],\n}\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\n\ndef build_model(name, num_classes=NUM_CLASSES):\n    if name == \"vgg19\":\n        m = tvm.vgg19(weights=tvm.VGG19_Weights.IMAGENET1K_V1)\n        in_f = m.classifier[6].in_features\n        m.classifier[6] = nn.Linear(in_f, num_classes)\n    elif name == \"resnet50\":\n        m = tvm.resnet50(weights=tvm.ResNet50_Weights.IMAGENET1K_V2)\n        in_f = m.fc.in_features\n        m.fc = nn.Linear(in_f, num_classes)\n    elif name == \"inception_v3\":\n        m = tvm.inception_v3(weights=tvm.Inception_V3_Weights.IMAGENET1K_V1, aux_logits=True)\n        m.fc = nn.Linear(m.fc.in_features, num_classes)\n        m.AuxLogits.fc = nn.Linear(m.AuxLogits.fc.in_features, num_classes)\n    elif name == \"xception\":\n        m = timm.create_model(\"legacy_xception\", pretrained=True, num_classes=num_classes)\n    elif name == \"efficientnet_b0\":\n        m = timm.create_model(\"efficientnet_b0\", pretrained=True, num_classes=num_classes)\n    elif name == \"densenet121\":\n        m = timm.create_model(\"densenet121\", pretrained=True, num_classes=num_classes)\n    else:\n        raise ValueError(name)\n    return m.to(DEVICE)\n\n\ndef make_classifier_transform(input_size):\n    return T.Compose([\n        T.ToPILImage(),\n        T.Resize((input_size, input_size)),\n        T.Grayscale(num_output_channels=3),  # repeat single channel -> RGB\n        T.ToTensor(),\n        T.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n    ])\n\n\nclass ArrayDataset(Dataset):\n    '''Wraps an in-memory (X, y) pair of uint8 grayscale images with a per-model transform.'''\n    def __init__(self, X, y, transform):\n        self.X, self.y, self.transform = X, y, transform\n\n    def __len__(self):\n        return len(self.y)\n\n    def __getitem__(self, idx):\n        img = self.transform(self.X[idx])\n        return img, int(self.y[idx])\n","metadata":{},"outputs":[],"execution_count":null},{"id":"431afc4c","cell_type":"markdown","source":"## 12. Classifier training & evaluation loop","metadata":{}},{"id":"99183ce0","cell_type":"code","source":"\ndef get_head_module(model, model_name):\n    '''Return the final classification-head module for a built model, so its\n    parameters can be kept trainable while the backbone is frozen during warmup.'''\n    if model_name == \"vgg19\":\n        return model.classifier[6]\n    if model_name in (\"resnet50\", \"inception_v3\"):\n        return model.fc\n    if model_name in (\"xception\", \"efficientnet_b0\", \"densenet121\"):\n        # all three are built via timm.create_model(...); get_classifier() is the\n        # generic timm accessor for the final linear layer regardless of architecture\n        return model.get_classifier()\n    raise ValueError(model_name)\n\n\ndef set_backbone_trainable(model, model_name, trainable):\n    '''Freeze/unfreeze every parameter except the classification head. Used for\n    staged fine-tuning: head-only warmup, then full-network fine-tuning.'''\n    for p in model.parameters():\n        p.requires_grad = trainable\n    for p in get_head_module(model, model_name).parameters():\n        p.requires_grad = True\n    if model_name == \"inception_v3\" and hasattr(model, \"AuxLogits\"):\n        for p in model.AuxLogits.fc.parameters():\n            p.requires_grad = True\n","metadata":{},"outputs":[],"execution_count":null},{"id":"080c37d6","cell_type":"markdown","source":"\n## 12b. Staged fine-tuning\n\nEvery backbone here is already fine-tuned end-to-end (no layers were frozen in the\nfirst version of this notebook) -- but training the whole pretrained network from\nthe very first batch, with a randomly-initialized head sitting on top, risks large\nearly gradients through the head corrupting the pretrained backbone weights before\nthe head has learned anything sensible. The standard fix, used here when\n`use_staged_finetuning` is on:\n\n1. **Warmup** (`finetune_warmup_epochs`): freeze the backbone, train only the new\n   classification head at `classifier_lr`.\n2. **Fine-tune**: unfreeze the whole network, continue training all layers together\n   at a lower `finetune_lr` for the remaining epochs.\n","metadata":{}},{"id":"74a7251c","cell_type":"code","source":"\ndef train_classifier(model_name, X_train, y_train, X_test, y_test, epochs, batch_size, lr):\n    input_size = MODEL_INPUT_SIZES[model_name]\n    transform = make_classifier_transform(input_size)\n\n    # --- carve a small stratified validation slice out of the training set (in-memory only,\n    # never touches X_test) so we can checkpoint the best epoch instead of trusting the last one ---\n    val_frac = IMPLEMENTATION_ASSUMPTIONS.get(\"val_frac\", 0.0)\n    if val_frac > 0 and len(np.unique(y_train)) > 1:\n        X_fit, X_val, y_fit, y_val = train_test_split(\n            X_train, y_train, test_size=val_frac, stratify=y_train, random_state=SEED)\n    else:\n        X_fit, y_fit = X_train, y_train\n        X_val, y_val = None, None\n\n    train_ds = ArrayDataset(X_fit, y_fit, transform)\n    test_ds = ArrayDataset(X_test, y_test, transform)\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True,\n                               num_workers=0, drop_last=False)\n    test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, num_workers=0)\n    val_loader = None\n    if X_val is not None:\n        val_ds = ArrayDataset(X_val, y_val, transform)\n        val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0)\n\n    model = build_model(model_name)\n\n    class_counts = np.bincount(y_fit, minlength=NUM_CLASSES).astype(np.float32)\n    class_weights = torch.tensor(class_counts.sum() / (NUM_CLASSES * np.maximum(class_counts, 1)),\n                                  dtype=torch.float32, device=DEVICE)\n    label_smoothing = IMPLEMENTATION_ASSUMPTIONS.get(\"label_smoothing\", 0.0)\n    criterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=label_smoothing)\n    weight_decay = IMPLEMENTATION_ASSUMPTIONS.get(\"classifier_weight_decay\", 0.0)\n\n    use_staged = IMPLEMENTATION_ASSUMPTIONS[\"use_staged_finetuning\"]\n    warmup_epochs = min(IMPLEMENTATION_ASSUMPTIONS[\"finetune_warmup_epochs\"], epochs) if use_staged else 0\n    finetune_lr = IMPLEMENTATION_ASSUMPTIONS[\"finetune_lr\"]\n\n    if use_staged and warmup_epochs > 0:\n        set_backbone_trainable(model, model_name, trainable=False)\n        optimizer = torch.optim.AdamW(\n            [p for p in model.parameters() if p.requires_grad], lr=lr, weight_decay=weight_decay)\n    else:\n        optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n\n    best_val_acc = -1.0\n    best_state = None\n    epochs_since_improve = 0\n    patience = IMPLEMENTATION_ASSUMPTIONS.get(\"early_stop_patience\", 0)\n    min_delta = IMPLEMENTATION_ASSUMPTIONS.get(\"early_stop_min_delta\", 0.0)\n\n    use_amp = IMPLEMENTATION_ASSUMPTIONS.get(\"use_amp\", False) and torch.cuda.is_available()\n    scaler = torch.cuda.amp.GradScaler(enabled=use_amp)\n\n    history = {\"epoch\": [], \"train_loss\": [], \"train_acc\": [], \"val_acc\": [], \"stage\": []}\n    for epoch in range(epochs):\n        if use_staged and epoch == warmup_epochs:\n            # switch from head-only warmup to full-network fine-tuning\n            set_backbone_trainable(model, model_name, trainable=True)\n            optimizer = torch.optim.AdamW(model.parameters(), lr=finetune_lr, weight_decay=weight_decay)\n\n        stage_name = \"warmup\" if (use_staged and epoch < warmup_epochs) else \"finetune\"\n\n        model.train()\n        total, correct, running_loss = 0, 0, 0.0\n        for xb, yb in train_loader:\n            xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n            optimizer.zero_grad()\n            with torch.cuda.amp.autocast(enabled=use_amp):\n                out = model(xb)\n                if model_name == \"inception_v3\" and isinstance(out, tuple):\n                    main_out, aux_out = out.logits, out.aux_logits\n                    loss = criterion(main_out, yb) + 0.4 * criterion(aux_out, yb)\n                    logits = main_out\n                else:\n                    logits = out\n                    loss = criterion(logits, yb)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            running_loss += loss.item() * xb.size(0)\n            preds = logits.argmax(dim=1)\n            correct += (preds == yb).sum().item()\n            total += xb.size(0)\n\n        # --- validation pass: track accuracy, checkpoint the best epoch's weights ---\n        val_acc = None\n        if val_loader is not None:\n            model.eval()\n            v_total, v_correct = 0, 0\n            with torch.no_grad():\n                for xb, yb in val_loader:\n                    xb, yb = xb.to(DEVICE), yb.to(DEVICE)\n                    with torch.cuda.amp.autocast(enabled=use_amp):\n                        out = model(xb)\n                    logits = out.logits if (model_name == \"inception_v3\" and hasattr(out, \"logits\")) else out\n                    preds = logits.argmax(dim=1)\n                    v_correct += (preds == yb).sum().item()\n                    v_total += xb.size(0)\n            val_acc = v_correct / max(v_total, 1)\n            if val_acc >= best_val_acc + min_delta:\n                best_val_acc = val_acc\n                best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n                epochs_since_improve = 0\n            else:\n                epochs_since_improve += 1\n\n        history[\"epoch\"].append(epoch)\n        history[\"train_loss\"].append(running_loss / max(total, 1))\n        history[\"train_acc\"].append(correct / max(total, 1))\n        history[\"val_acc\"].append(val_acc)\n        history[\"stage\"].append(stage_name)\n\n        # --- early stopping: only once we're past warmup (warmup accuracy is noisy/low by\n        # design since only the head is training) and only if a patience value is set ---\n        past_warmup = (not use_staged) or (epoch >= warmup_epochs)\n        if patience and past_warmup and val_loader is not None and epochs_since_improve >= patience:\n            print(f\"    early stop at epoch {epoch} (no val improvement for {patience} epochs)\")\n            break\n\n    # restore the best-validation-epoch weights before returning (falls back to last epoch\n    # if no validation slice was available)\n    if best_state is not None:\n        model.load_state_dict(best_state)\n\n    return model, history, test_loader\n\n\n@torch.no_grad()\ndef evaluate_classifier(model, model_name, test_loader):\n    model.eval()\n    all_logits, all_true = [], []\n    for xb, yb in test_loader:\n        xb = xb.to(DEVICE)\n        out = model(xb)\n        logits = out.logits if (model_name == \"inception_v3\" and hasattr(out, \"logits\")) else out\n        all_logits.append(logits.cpu().numpy())\n        all_true.append(yb.numpy())\n    logits = np.concatenate(all_logits, axis=0)\n    y_true = np.concatenate(all_true, axis=0)\n    probs = torch.softmax(torch.from_numpy(logits), dim=1).numpy()\n    y_pred = probs.argmax(axis=1)\n\n    accuracy = accuracy_score(y_true, y_pred)\n    precision, recall, f1, _ = precision_recall_fscore_support(\n        y_true, y_pred, labels=np.arange(NUM_CLASSES), zero_division=0)\n\n    per_class_rows = []\n    for class_id, class_name in enumerate(CLASSES):\n        binary_truth = (y_true == class_id).astype(int)\n        fpr, tpr, _ = roc_curve(binary_truth, probs[:, class_id])\n        class_auc = auc(fpr, tpr)\n        per_class_rows.append({\n            \"class\": class_name,\n            \"precision\": precision[class_id],\n            \"recall\": recall[class_id],\n            \"f1\": f1[class_id],\n            \"auc\": class_auc,\n        })\n\n    cm = confusion_matrix(y_true, y_pred, labels=np.arange(NUM_CLASSES))\n    return {\n        \"accuracy\": accuracy,\n        \"per_class\": per_class_rows,\n        \"confusion_matrix\": cm,\n        \"y_true\": y_true,\n        \"y_pred\": y_pred,\n        \"probs\": probs,\n    }\n","metadata":{},"outputs":[],"execution_count":null},{"id":"26490184","cell_type":"markdown","source":"\n## 13. Run the experiment grid\n\nEvery model architecture is trained on every point of the classic curve and every\npoint of the synthetic curve, then evaluated on the **same held-out real test set**\n(never augmented, never seen by any GAN). Results feed the Table-I-style summary and\nthe total-accuracy-vs-samples-per-class plot (Fig. 6.13 in the repo report).\n","metadata":{}},{"id":"4dc582af","cell_type":"code","source":"\nMODEL_NAMES = [\"vgg19\", \"inception_v3\", \"xception\", \"resnet50\", \"densenet121\"]\n\nif RECOVERY_MODE:\n    print(\"RECOVERY_MODE on -- skipping the full 32-combo grid (section 13b handles \"\n          \"reloading existing results + retraining only the best group per model).\")\n\nelse:\n    RESULTS_CSV = WORK_DIR / \"classifier_results.csv\"\n    CURVE_CSV = WORK_DIR / \"accuracy_curves.csv\"\n\n    # --- resume support: if this cell was already run partway through and the kernel/session\n    # died, reload whatever finished last time so we don't retrain combos we already have ---\n    if RESULTS_CSV.exists() and CURVE_CSV.exists():\n        results_df = pd.read_csv(RESULTS_CSV)\n        curve_df = pd.read_csv(CURVE_CSV)\n        results_rows = results_df.to_dict(\"records\")\n        curve_rows = curve_df.to_dict(\"records\")\n        completed = set(zip(curve_df[\"model\"], curve_df[\"curve\"], curve_df[\"samples_per_class\"]))\n        print(f\"Resuming: found {len(completed)} already-completed (model, curve, n) combos \"\n              f\"in {RESULTS_CSV.name} / {CURVE_CSV.name}. These will be skipped.\")\n    else:\n        results_rows = []\n        curve_rows = []\n        completed = set()\n\n    all_eval_results = {}  # (model_name, stage) -> full metrics dict, incl. confusion_matrix\n                           # NOTE: only populated for combos (re)trained THIS session -- if you\n                           # resumed after a crash, combos loaded from CSV above won't have a\n                           # confusion matrix here (their accuracy numbers are still in curve_df).\n\n    def save_progress():\n        pd.DataFrame(results_rows).to_csv(RESULTS_CSV, index=False)\n        pd.DataFrame(curve_rows).to_csv(CURVE_CSV, index=False)\n\n    if DATASET_ROOT is not None:\n        for model_name in MODEL_NAMES:\n            print(f\"\\n##### {model_name} #####\")\n\n            # --- classic curve ---\n            for n, (X_tr, y_tr) in classic_curve_datasets.items():\n                if (model_name, \"classic\", n) in completed:\n                    print(f\"  classic C={n}: already done, skipping\")\n                    continue\n                model, hist, test_loader = train_classifier(\n                    model_name, X_tr, y_tr, X_test, y_test,\n                    epochs=IMPLEMENTATION_ASSUMPTIONS[\"classifier_epochs\"],\n                    batch_size=IMPLEMENTATION_ASSUMPTIONS[\"classifier_batch_size\"],\n                    lr=IMPLEMENTATION_ASSUMPTIONS[\"classifier_lr\"],\n                )\n                metrics = evaluate_classifier(model, model_name, test_loader)\n                stage_key = f\"classic_C{n}\"\n                all_eval_results[(model_name, stage_key)] = metrics\n                print(f\"  classic C={n}: total_accuracy={metrics['accuracy']:.4f}\")\n                curve_rows.append({\"model\": model_name, \"curve\": \"classic\",\n                                    \"samples_per_class\": n, \"total_accuracy\": metrics[\"accuracy\"]})\n                for row in metrics[\"per_class\"]:\n                    results_rows.append({\"model\": model_name, \"stage\": stage_key, **row})\n                save_progress()  # write after EVERY group finishes, not just at the very end\n                del model, test_loader\n                gc.collect()\n                if torch.cuda.is_available():\n                    torch.cuda.empty_cache()\n\n            # --- synthetic curve ---\n            for n, (X_tr, y_tr) in synthetic_curve_datasets.items():\n                if (model_name, \"synthetic\", n) in completed:\n                    print(f\"  synthetic G={n}: already done, skipping\")\n                    continue\n                model, hist, test_loader = train_classifier(\n                    model_name, X_tr, y_tr, X_test, y_test,\n                    epochs=IMPLEMENTATION_ASSUMPTIONS[\"classifier_epochs\"],\n                    batch_size=IMPLEMENTATION_ASSUMPTIONS[\"classifier_batch_size\"],\n                    lr=IMPLEMENTATION_ASSUMPTIONS[\"classifier_lr\"],\n                )\n                metrics = evaluate_classifier(model, model_name, test_loader)\n                stage_key = f\"synthetic_G{n}\"\n                all_eval_results[(model_name, stage_key)] = metrics\n                print(f\"  synthetic G={n}: total_accuracy={metrics['accuracy']:.4f}\")\n                curve_rows.append({\"model\": model_name, \"curve\": \"synthetic\",\n                                    \"samples_per_class\": n, \"total_accuracy\": metrics[\"accuracy\"]})\n                for row in metrics[\"per_class\"]:\n                    results_rows.append({\"model\": model_name, \"stage\": stage_key, **row})\n                save_progress()  # write after EVERY group finishes, not just at the very end\n                del model, test_loader\n                gc.collect()\n                if torch.cuda.is_available():\n                    torch.cuda.empty_cache()\n\n        # rebuild the DataFrame variables from the final row lists so downstream cells\n        # (section 14 onward) have them even on a fresh run with no prior resume\n        results_df = pd.DataFrame(results_rows)\n        curve_df = pd.DataFrame(curve_rows)\n        save_progress()\n        print(\"\\nSaved classifier_results.csv and accuracy_curves.csv\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"eec6f907","cell_type":"markdown","source":"## 13b. Recovery path — controlled by `RECOVERY_MODE` in the config cell\n\nThis notebook now supports two modes via a single flag set in section 1 (Configuration):\n\n- **`RECOVERY_MODE = True`** (current default): skips GAN training (section 7) and the full 32-combo grid (section 13). Instead: loads previously-saved synthetic image pools, reloads previous CSV results, and retrains only the single best-performing group per model (4 runs total) so confusion matrices and Grad-CAM have real data to work with. Use this after a previous run already produced `accuracy_curves.csv` / `classifier_results.csv` / `synthetic_pool_*.npy` (attached as an input dataset, path set via `RECOVERY_SOURCE_DIR` in section 1).\n- **`RECOVERY_MODE = False`**: a genuine first-time run. Trains all GANs from scratch, runs the full 32-combo grid, and this cell does nothing (section 13 already populated everything it needs).\n\n**You can just click Run All** with the current settings — every cell below checks the flag itself and skips or runs the right thing automatically.\n","metadata":{}},{"id":"9987dd91","cell_type":"code","source":"if RECOVERY_MODE:\n    # --- 1. Rebuild the exact split from the saved manifest (no re-shuffling) ---\n    split_manifest = pd.read_csv(WORK_DIR / \"split_manifest.csv\")\n    SPLIT = {\"train\": {}, \"test\": {}}\n    for split_name in [\"train\", \"test\"]:\n        for cls in CLASSES:\n            paths = split_manifest.loc[\n                (split_manifest[\"split\"] == split_name) & (split_manifest[\"class\"] == cls), \"path\"\n            ].tolist()\n            SPLIT[split_name][cls] = [Path(p) for p in paths]\n    print(\"Rebuilt SPLIT from split_manifest.csv:\",\n          {k: {c: len(v) for c, v in d.items()} for k, d in SPLIT.items()})\n\n    # --- 2. Reload real train/test images (fast: only real_train/test_per_class images) ---\n    REAL_TRAIN_IMAGES = real_only_images()\n    REAL_TEST_IMAGES = real_test_images()\n    X_test, y_test = assemble_dataset(REAL_TEST_IMAGES)\n\n    # --- 3. Reload the already-generated synthetic pools instead of retraining any GAN ---\n    def _load_synth_pool(cls):\n        for candidate in (cls, cls[:6]):  # handles any filename truncation (e.g. PNEUMONIA -> PNEUMO)\n            p = WORK_DIR / f\"synthetic_pool_{candidate}.npy\"\n            if p.exists():\n                return np.load(p)\n        raise FileNotFoundError(f\"no synthetic_pool_*.npy found for class {cls}\")\n\n    synthetic_pools = {cls: _load_synth_pool(cls) for cls in CLASSES}\n    print(\"Reloaded synthetic pools:\", {c: v.shape for c, v in synthetic_pools.items()})\n\n    # --- 4. Reload the full 32-combo results as-is -- sections 14/15/15b/17 need nothing else ---\n    results_df = pd.read_csv(WORK_DIR / \"classifier_results.csv\")\n    curve_df = pd.read_csv(WORK_DIR / \"accuracy_curves.csv\")\n    results_rows = results_df.to_dict(\"records\")\n    curve_rows = curve_df.to_dict(\"records\")\n\n    # --- 5. Rebuild classic-augmented base (deterministic given the same SEED) ---\n    classic_groups = PAPER_PROTOCOL[\"classic_groups_per_class\"]\n    best_classic_n = max(classic_groups)\n    best_classic_imgs = {\n        cls: build_classic_augmented_set(SPLIT[\"train\"][cls], best_classic_n)\n        for cls in CLASSES\n    }\n\n    # --- 6. For each model, find its single best (curve, n) from the already-saved results,\n    #     rebuild ONLY that one training set, retrain ONLY that one run, and this time save\n    #     everything needed for sections 15c/16: confusion matrix + the actual trained weights ---\n    MODEL_WEIGHTS_DIR = WORK_DIR / \"best_classifier_weights\"\n    MODEL_WEIGHTS_DIR.mkdir(exist_ok=True)\n    all_eval_results = {}\n    best_trained_models = {}\n\n    for model_name in MODEL_NAMES:\n        sub = curve_df[curve_df[\"model\"] == model_name]\n        if not len(sub):\n            continue\n        best = sub.loc[sub[\"total_accuracy\"].idxmax()]\n        n = int(best[\"samples_per_class\"])\n        curve_name = best[\"curve\"]\n        stage_key = f\"classic_C{n}\" if curve_name == \"classic\" else f\"synthetic_G{n}\"\n        print(f\"\\n{model_name}: retraining only its best group -> {stage_key} \"\n              f\"(previously scored {best['total_accuracy']:.4f})\")\n\n        combo = {}\n        for cls in CLASSES:\n            if curve_name == \"classic\":\n                classic_imgs = build_classic_augmented_set(SPLIT[\"train\"][cls], n)\n                combo[cls] = np.concatenate([REAL_TRAIN_IMAGES[cls], classic_imgs], axis=0)\n            else:\n                synth_imgs = sample_from_pool(synthetic_pools[cls], n)\n                combo[cls] = np.concatenate(\n                    [REAL_TRAIN_IMAGES[cls], best_classic_imgs[cls], synth_imgs], axis=0\n                )\n        X_tr, y_tr = assemble_dataset(combo)\n\n        model, hist, test_loader = train_classifier(\n            model_name, X_tr, y_tr, X_test, y_test,\n            epochs=IMPLEMENTATION_ASSUMPTIONS[\"classifier_epochs\"],\n            batch_size=IMPLEMENTATION_ASSUMPTIONS[\"classifier_batch_size\"],\n            lr=IMPLEMENTATION_ASSUMPTIONS[\"classifier_lr\"],\n        )\n        metrics = evaluate_classifier(model, model_name, test_loader)\n        all_eval_results[(model_name, stage_key)] = metrics\n        print(f\"  re-verified accuracy: {metrics['accuracy']:.4f}\")\n\n        weight_path = MODEL_WEIGHTS_DIR / f\"{model_name}_{stage_key}.pt\"\n        torch.save(model.state_dict(), weight_path)\n        best_trained_models[model_name] = model  # keep in memory for the Grad-CAM cell below\n\n        del test_loader\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    print(\"\\nRecovery complete: all_eval_results has real confusion matrices, \"\n          \"best_trained_models has real fine-tuned weights, saved to \", MODEL_WEIGHTS_DIR)\n\nelse:\n    print(\"RECOVERY_MODE off -- section 13 already produced full results and all_eval_results this session, nothing to recover.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"67dcc8cb","cell_type":"markdown","source":"## 14. Total-accuracy-vs-samples-per-class curves (Fig. 6.13 style), per model","metadata":{}},{"id":"9aaba74c","cell_type":"code","source":"\nif DATASET_ROOT is not None and len(curve_rows):\n    for model_name in MODEL_NAMES:\n        sub = curve_df[curve_df[\"model\"] == model_name]\n        plt.figure(figsize=(5.5, 4))\n        for curve_name, marker in [(\"classic\", \"s-\"), (\"synthetic\", \"s-\")]:\n            cs = sub[sub[\"curve\"] == curve_name].sort_values(\"samples_per_class\")\n            if len(cs):\n                plt.plot(cs[\"samples_per_class\"], cs[\"total_accuracy\"], marker, label=curve_name)\n        plt.xlabel(\"Samples per class\")\n        plt.ylabel(\"Total accuracy\")\n        plt.title(f\"{model_name}: classic vs. synthetic augmentation\")\n        plt.legend()\n        plt.grid(alpha=0.3)\n        plt.tight_layout()\n        plt.show(); plt.close()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"4bc7db49","cell_type":"markdown","source":"## 15. Table-I-style results summary (best group per model)","metadata":{}},{"id":"f2fe949d","cell_type":"code","source":"\nif DATASET_ROOT is not None and len(results_rows):\n    best_rows = []\n    for model_name in MODEL_NAMES:\n        model_curve = curve_df[curve_df[\"model\"] == model_name]\n        if not len(model_curve):\n            continue\n        best = model_curve.loc[model_curve[\"total_accuracy\"].idxmax()]\n        stage_key = (f\"classic_C{best['samples_per_class']}\" if best[\"curve\"] == \"classic\"\n                     else f\"synthetic_G{best['samples_per_class']}\")\n        stage_rows = results_df[(results_df[\"model\"] == model_name) &\n                                 (results_df[\"stage\"] == stage_key)]\n        for _, r in stage_rows.iterrows():\n            best_rows.append({\n                \"Model\": model_name, \"Best stage\": stage_key,\n                \"Total accuracy\": round(best[\"total_accuracy\"], 4),\n                \"Class\": r[\"class\"], \"Precision\": round(r[\"precision\"], 4),\n                \"Recall\": round(r[\"recall\"], 4), \"F1\": round(r[\"f1\"], 4),\n                \"AUC\": round(r[\"auc\"], 4),\n            })\n    display(pd.DataFrame(best_rows))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a09791c7","cell_type":"markdown","source":"\n## 15b. Full accuracy table across every group (every model x every augmentation point)\n\nThe table above only shows each model's single best-performing group, matching how\nthe OT paper reports one configuration per model. This table shows **total accuracy\nfor every group** so the full picture -- including any point where accuracy goes\ndown as samples increase -- is visible, not just the best point.\n","metadata":{}},{"id":"e212e589","cell_type":"code","source":"\nif DATASET_ROOT is not None and len(curve_rows):\n    full_table = curve_df.pivot_table(\n        index=\"model\", columns=[\"curve\", \"samples_per_class\"], values=\"total_accuracy\"\n    ).round(4)\n    display(full_table)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"049f9f5a","cell_type":"markdown","source":"\n## 15c. Confusion matrix per model (best group), matching the OT paper's Table I + Fig. 4\n\nOne confusion matrix per architecture, computed on the held-out real test set at\neach model's best-performing group.\n","metadata":{}},{"id":"5ceb95f5","cell_type":"code","source":"\nif DATASET_ROOT is not None and len(all_eval_results):\n    fig, axes = plt.subplots(1, len(MODEL_NAMES), figsize=(4.2 * len(MODEL_NAMES), 4))\n    if len(MODEL_NAMES) == 1:\n        axes = [axes]\n    for ax, model_name in zip(axes, MODEL_NAMES):\n        model_curve = curve_df[curve_df[\"model\"] == model_name]\n        if not len(model_curve):\n            ax.axis(\"off\")\n            continue\n        best = model_curve.loc[model_curve[\"total_accuracy\"].idxmax()]\n        stage_key = (f\"classic_C{best['samples_per_class']}\" if best[\"curve\"] == \"classic\"\n                     else f\"synthetic_G{best['samples_per_class']}\")\n        cm = all_eval_results[(model_name, stage_key)][\"confusion_matrix\"]\n        cm_norm = cm.astype(float) / cm.sum(axis=1, keepdims=True).clip(min=1)\n\n        im = ax.imshow(cm_norm, cmap=\"Blues\", vmin=0, vmax=1)\n        ax.set_xticks(range(NUM_CLASSES)); ax.set_xticklabels(CLASSES, rotation=45, ha=\"right\")\n        ax.set_yticks(range(NUM_CLASSES)); ax.set_yticklabels(CLASSES)\n        ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\" if model_name == MODEL_NAMES[0] else \"\")\n        for i in range(NUM_CLASSES):\n            for j in range(NUM_CLASSES):\n                ax.text(j, i, f\"{cm_norm[i,j]:.2f}\", ha=\"center\", va=\"center\",\n                        color=\"white\" if cm_norm[i, j] > 0.5 else \"black\", fontsize=9)\n        ax.set_title(f\"{model_name}\\n{stage_key}, acc={best['total_accuracy']:.3f}\")\n    plt.tight_layout()\n    plt.show(); plt.close()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"e76cccdd","cell_type":"markdown","source":"\n## 16. Grad-CAM explainability (optional, matches OT paper Section IV-B)\n\nGeneric implementation via forward/backward hooks on the last convolutional feature\nmap, so the same function works across VGG19/ResNet50/InceptionV3/Xception.\n","metadata":{}},{"id":"8ca41868","cell_type":"code","source":"\ndef find_last_conv_module_name(model):\n    '''Auto-detect the last nn.Conv2d module by walking every submodule in forward\n    (creation) order and keeping the last match. Replaces a hardcoded per-architecture\n    name lookup -- the previous version's own comment flagged that timm could change\n    internal naming and break it for any single architecture; auto-detection sidesteps\n    that entirely and works for all five backbones here, including the two new ones\n    whose exact internal layer names weren't verified against this timm version.'''\n    last_name = None\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d):\n            last_name = name\n    if last_name is None:\n        raise ValueError(\"no Conv2d module found in this model\")\n    return last_name\n\n\ndef _disable_inplace_ops(model):\n    '''VGG19 (and some other torchvision backbones) use inplace=True ReLU by default,\n    which conflicts with register_full_backward_hook: \"Output 0 of\n    BackwardHookFunctionBackward is a view and is being modified inplace.\" Standard\n    fix (used by e.g. the pytorch-grad-cam library): switch every inplace-capable\n    module to inplace=False before running Grad-CAM. Purely a Grad-CAM-time setting,\n    doesn't affect the model's trained weights or normal forward/backward behavior.'''\n    for module in model.modules():\n        if hasattr(module, \"inplace\"):\n            module.inplace = False\n\n\ndef gradcam(model, model_name, image_tensor, target_class=None):\n    model.eval()\n    _disable_inplace_ops(model)\n    activations, gradients = {}, {}\n\n    layer_name = find_last_conv_module_name(model)\n    layer = dict(model.named_modules())[layer_name]\n\n    def fwd_hook(module, inp, out):\n        activations[\"value\"] = out.detach()\n\n    def bwd_hook(module, grad_in, grad_out):\n        gradients[\"value\"] = grad_out[0].detach()\n\n    h1 = layer.register_forward_hook(fwd_hook)\n    h2 = layer.register_full_backward_hook(bwd_hook)\n\n    x = image_tensor.unsqueeze(0).to(DEVICE).requires_grad_(True)\n    out = model(x)\n    logits = out.logits if (model_name == \"inception_v3\" and hasattr(out, \"logits\")) else out\n    if target_class is None:\n        target_class = logits.argmax(dim=1).item()\n    score = logits[0, target_class]\n    model.zero_grad()\n    score.backward()\n\n    h1.remove(); h2.remove()\n\n    acts = activations[\"value\"][0]     # (C, h, w)\n    grads = gradients[\"value\"][0]      # (C, h, w)\n    weights = grads.mean(dim=(1, 2))   # (C,)\n    cam = F.relu((weights[:, None, None] * acts).sum(dim=0))\n    cam = cam / (cam.max() + 1e-8)\n    return cam.cpu().numpy(), target_class\n\n\nif DATASET_ROOT is not None and len(results_rows):\n    example_model_name = MODEL_NAMES[0]\n    input_size = MODEL_INPUT_SIZES[example_model_name]\n    transform = make_classifier_transform(input_size)\n    if \"best_trained_models\" in dir() and example_model_name in best_trained_models:\n        demo_model = best_trained_models[example_model_name]\n    else:\n        print(\"WARNING: no fine-tuned weights available in memory for \"\n              f\"'{example_model_name}' -- Grad-CAM below uses an UNTRAINED classifier head \"\n              \"and is not meaningful. Run section 13b first to get real trained weights.\")\n        demo_model = build_model(example_model_name)\n\n    fig, axes = plt.subplots(1, len(CLASSES), figsize=(4 * len(CLASSES), 4))\n    for i, cls in enumerate(CLASSES):\n        raw_img = REAL_TEST_IMAGES[cls][0]\n        tensor = transform(raw_img)\n        try:\n            cam, pred_class = gradcam(demo_model, example_model_name, tensor)\n            axes[i].imshow(tensor[0].cpu().numpy(), cmap=\"gray\")\n            axes[i].imshow(cam, cmap=\"jet\", alpha=0.5,\n                            extent=(0, tensor.shape[-1], tensor.shape[-1], 0))\n            axes[i].set_title(f\"{cls} (pred: {CLASSES[pred_class]})\")\n        except Exception as e:\n            axes[i].set_title(f\"{cls}: gradcam failed ({type(e).__name__})\")\n        axes[i].axis(\"off\")\n    plt.suptitle(f\"Grad-CAM — {example_model_name}\")\n    plt.tight_layout()\n    plt.show(); plt.close()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b4266c79","cell_type":"markdown","source":"## 17. Leakage audit — must pass before trusting any result above","metadata":{}},{"id":"4603b33b","cell_type":"code","source":"\nif DATASET_ROOT is not None:\n    train_paths_check = set(split_manifest.loc[split_manifest[\"split\"] == \"train\", \"path\"])\n    test_paths_check = set(split_manifest.loc[split_manifest[\"split\"] == \"test\", \"path\"])\n\n    assert train_paths_check.isdisjoint(test_paths_check), \\\n        \"LEAKAGE: a file path appears in both train and test manifests.\"\n\n    for cls in CLASSES:\n        assert len(SPLIT[\"test\"][cls]) == IMPLEMENTATION_ASSUMPTIONS[\"real_test_per_class\"], \\\n            f\"LEAKAGE CHECK FAILED: test set size for {cls} changed unexpectedly.\"\n\n    # The GAN was only ever handed SPLIT[\"train\"][cls] (see make_gan_loader) -- never\n    # SPLIT[\"test\"][cls] -- so synthetic images cannot be derived from test images.\n    # The classic augmentation function was only ever called with SPLIT[\"train\"][cls]\n    # paths as well. X_test/y_test were assembled purely from REAL_TEST_IMAGES.\n    print(\"Leakage audit passed: train/test paths are disjoint, test set untouched by \"\n          \"GAN training and classic augmentation, and evaluation used only the \"\n          \"held-out real test set.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"dfeb548c","cell_type":"markdown","source":"\n## 18. Notes on deviations & how to read the results\n\n- The OT paper reports VGG19/InceptionV3/Xception accuracies of 96.94% / 95.35% /\n  92.62% on a **4-class** task with 1600 synthetic images generated from the full\n  Sinkhorn-regularized GAN described in its Eq. 2-6. This notebook now matches the\n  **4-class scope** using the COVID-19 Radiography Database, but still uses this\n  repo's much smaller 64x64 GAN architecture and a **combination** of Wasserstein +\n  Sinkhorn loss that neither source specifies in exact form — so matching those\n  numbers exactly is still not expected, though the 4-class scope removes one\n  difference from the earlier 3-class run.\n- `IMPLEMENTATION_ASSUMPTIONS[\"quick_debug\"]` controls fast smoke-test mode vs. a\n  full paper-scale run. This is currently set for a full run, tuned for a 5-6h\n  session: 5 backbones × 2 classic (`{200,1000}`) + 2 synthetic (`{100,900}`)\n  groups = 20 classifier training runs (down from 48 in the 6-backbone version,\n  which hit Kaggle's session limit at ~11.4h and had to be cancelled), plus 4\n  per-class GANs at 4000 iterations each (down from 6000). The two group values\n  test the range endpoints rather than the middle, based on evidence from a\n  separate run that several models' real best point sits at an extreme.\n- Grad-CAM disables inplace ops on the model before hooking it, avoiding a\n  `RuntimeError` that VGG19's default `inplace=True` ReLUs trigger with\n  `register_full_backward_hook` — a known PyTorch interaction, fixed without\n  changing anything about how the model was trained or evaluated elsewhere.\n- ResNet50 and DenseNet121 have no baseline in the source paper; they're included\n  as additional comparison points. DenseNet121 in particular is the architecture\n  CheXNet used, so it's a meaningful domain-relevant addition rather than an\n  arbitrary one.\n- **EfficientNet-B0 was tried and dropped.** In the cancelled 6-backbone run it\n  scored 70-76% accuracy against 85-92% for every other model — a real\n  underperformance, not just a slower architecture — so cutting it saves time\n  without giving up a competitive result. Its code path in `build_model`/\n  `get_head_module`/`MODEL_INPUT_SIZES` (`efficientnet_b0_input`) is left in place\n  if you want to re-add and debug it later — append `\"efficientnet_b0\"` back into\n  `MODEL_NAMES` in cell 33.\n- VGG19 is included alongside the other four backbones this run (`MODEL_NAMES` in\n  cell 33 lists all five) — it's the paper's own best-performing model, so keeping\n  it in the comparison is the right default.\n","metadata":{}}]}