{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"9fa8d3b3","cell_type":"markdown","source":"# OT-GAN medical X-ray reproduction — paper-reported protocol\n\nThis notebook implements every concrete parameter and processing step reported in:\n\n**“Optimal Transport Theory-based GAN for Medical Image Augmentation and Classification” (ICCIT 2025).**\n\n## Reported settings implemented literally\n\n- Four classes: **Normal, COVID-19, Viral Pneumonia, Lung Opacity**\n- Two Kaggle sources: **Chest X-ray (COVID-19 & Pneumonia)** and **RSNA/Lung Opacity**\n- **25 real images per class**, 100 real images total\n- **1,600 generated images**, implemented as 400 per class\n- Stratified random sampling without replacement: **80% training / 20% testing**\n- Z-normalization\n- Three binary models: Normal versus COVID-19\n- Three multiclass models: four-class classification\n- Pretrained **VGG19**, **InceptionV3**, and **Xception**\n- Input sizes: VGG19/InceptionV3 **224×224**, Xception **299×299**\n- Weighted categorical loss\n- Adam with TensorFlow/Keras default parameters\n- Batch size **16**\n- Training for **500 epochs**\n- VGG19/InceptionV3/Xception heads containing the pooling, flattening, dropout where reported, and dense classification components\n- Xception average-pooling size **4×4**\n- WGAN critic with linear output\n- Wasserstein generator and critic objectives\n- Critic weight clipping to **[-0.01, 0.01] after every critic batch**\n- Entropy-regularized Sinkhorn OT term\n- Generator progression **64×64 → 512×512** with Conv2DTranspose, kernel 4, stride 2, and LeakyReLU\n- Critic progression **512×512 → 64×64** with stride-2 convolution and LeakyReLU\n- ROC, confusion matrix, accuracy, precision/PPV, recall/sensitivity, F1, AUC\n- Grad-CAM heatmaps and superimposed images\n\n## Parameters the paper does not disclose\n\nThe paper does not report the GAN latent dimension, GAN epoch count, number of critic updates, GAN learning rate, Sinkhorn epsilon, Sinkhorn coefficient, Sinkhorn iterations, validation fraction, random seed, channel widths, dropout rate, or whether the GAN was conditional or class-specific. Those values are isolated in `UNREPORTED_ASSUMPTIONS`; they are not presented as paper parameters.\n\nThe paper says the discriminator downsamples with `Conv2DTranspose`. A transpose convolution with stride 2 upsamples rather than downsamples, so the executable implementation uses `Conv2D` with the reported kernel, stride, and LeakyReLU while preserving the stated 512→64 geometry.","metadata":{}},{"id":"bb954e39","cell_type":"code","source":"import os\nimport gc\nimport glob\nimport json\nimport random\nimport subprocess\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_recall_fscore_support,\n    roc_curve,\n    auc,\n    confusion_matrix,\n    classification_report,\n)\n\ntry:\n    import pydicom\nexcept ImportError:\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"pydicom\", \"-q\"], check=True)\n    import pydicom\n\nprint(\"TensorFlow:\", tf.__version__)\nprint(\"GPUs:\", tf.config.list_physical_devices(\"GPU\"))\n\nfor gpu in tf.config.list_physical_devices(\"GPU\"):\n    try:\n        tf.config.experimental.set_memory_growth(gpu, True)\n    except Exception:\n        pass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:13:12.155339Z","iopub.execute_input":"2026-07-10T17:13:12.155609Z","iopub.status.idle":"2026-07-10T17:13:31.062740Z","shell.execute_reply.started":"2026-07-10T17:13:12.155574Z","shell.execute_reply":"2026-07-10T17:13:31.061985Z"}},"outputs":[],"execution_count":null},{"id":"33eb6138","cell_type":"markdown","source":"## 1. Protocol configuration","metadata":{}},{"id":"3ebe459f","cell_type":"code","source":"# Values stated explicitly in the paper.\nPAPER_PROTOCOL = {\n    \"classes\": [\"Normal\", \"COVID-19\", \"Viral Pneumonia\", \"Lung Opacity\"],\n    \"real_images_per_class\": 25,\n    \"generated_images_total\": 1600,\n    \"generated_images_per_class\": 400,\n    \"train_fraction\": 0.80,\n    \"test_fraction\": 0.20,\n    \"classifier_batch_size\": 16,\n    \"classifier_epochs\": 500,\n    \"vgg19_input\": 224,\n    \"inceptionv3_input\": 224,\n    \"xception_input\": 299,\n    \"generator_start_size\": 64,\n    \"gan_image_size\": 512,\n    \"gan_kernel_size\": 4,\n    \"gan_stride\": 2,\n    \"critic_weight_clip\": 1e-2,\n    \"xception_pool_size\": 4,\n    \"optimizer\": \"Adam(default parameters)\",\n    \"normalization\": \"feature-wise Z-normalization\",\n}\n\n# Required only because the paper does not disclose them.\n# Change these only when performing a sensitivity study; changing PAPER_PROTOCOL\n# means the run no longer follows the stated protocol.\nUNREPORTED_ASSUMPTIONS = {\n    \"random_seed\": 42,\n    \"validation_fraction_within_training_pool\": 0.20,\n    \"gan_latent_dim\": 100,\n    \"gan_epochs\": 500,\n    \"gan_batch_size\": 16,\n    \"critic_updates_per_generator\": 5,\n    \"sinkhorn_epsilon\": 0.10,\n    \"sinkhorn_weight\": 1.00,\n    \"sinkhorn_iterations\": 50,\n    \"inception_dropout_rate\": 0.50,\n    \"generator_start_channels\": 32,\n}\n\nRUN_CONTROL = {\n    # Full paper workflow: initial real-only training, GAN augmentation, retraining.\n    \"run_real_only_baseline\": True,\n    \"run_binary_models\": True,\n    \"run_multiclass_models\": True,\n    \"run_gan_training\": True,\n    \"run_augmented_models\": True,\n    \"models\": [\"VGG19\", \"InceptionV3\", \"Xception\"],\n    # Keep False for the paper protocol.\n    \"quick_debug\": False,\n}\n\nif RUN_CONTROL[\"quick_debug\"]:\n    print(\"WARNING: QUICK DEBUG MODE IS ACTIVE; THIS IS NOT THE PAPER PROTOCOL.\")\n    CLASSIFIER_EPOCHS = 2\n    GAN_EPOCHS = 2\nelse:\n    CLASSIFIER_EPOCHS = PAPER_PROTOCOL[\"classifier_epochs\"]\n    GAN_EPOCHS = UNREPORTED_ASSUMPTIONS[\"gan_epochs\"]\n\nSEED = UNREPORTED_ASSUMPTIONS[\"random_seed\"]\nos.environ[\"PYTHONHASHSEED\"] = str(SEED)\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntf.random.set_seed(SEED)\n\nCLASSES = PAPER_PROTOCOL[\"classes\"]\nNUM_CLASSES = len(CLASSES)\nCLASS_TO_ID = {name: i for i, name in enumerate(CLASSES)}\n\nWORK_DIR = Path(\"/kaggle/working/otgan_paper_reproduction\")\nWORK_DIR.mkdir(parents=True, exist_ok=True)\n(WORK_DIR / \"models\").mkdir(exist_ok=True)\n(WORK_DIR / \"histories\").mkdir(exist_ok=True)\n(WORK_DIR / \"predictions\").mkdir(exist_ok=True)\n(WORK_DIR / \"gan_cache\").mkdir(exist_ok=True)\n\ndisplay(pd.Series(PAPER_PROTOCOL, name=\"paper-reported value\").to_frame())\ndisplay(pd.Series(UNREPORTED_ASSUMPTIONS, name=\"necessary assumption\").to_frame())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:13:31.064249Z","iopub.execute_input":"2026-07-10T17:13:31.064813Z","iopub.status.idle":"2026-07-10T17:13:31.111312Z","shell.execute_reply.started":"2026-07-10T17:13:31.064787Z","shell.execute_reply":"2026-07-10T17:13:31.110516Z"}},"outputs":[],"execution_count":null},{"id":"faa73d3d","cell_type":"code","source":"def assert_paper_protocol():\n    assert PAPER_PROTOCOL[\"real_images_per_class\"] == 25\n    assert PAPER_PROTOCOL[\"generated_images_total\"] == 1600\n    assert PAPER_PROTOCOL[\"generated_images_per_class\"] * 4 == 1600\n    assert PAPER_PROTOCOL[\"train_fraction\"] == 0.80\n    assert PAPER_PROTOCOL[\"test_fraction\"] == 0.20\n    assert PAPER_PROTOCOL[\"classifier_batch_size\"] == 16\n    assert PAPER_PROTOCOL[\"classifier_epochs\"] == 500\n    assert PAPER_PROTOCOL[\"vgg19_input\"] == 224\n    assert PAPER_PROTOCOL[\"inceptionv3_input\"] == 224\n    assert PAPER_PROTOCOL[\"xception_input\"] == 299\n    assert PAPER_PROTOCOL[\"generator_start_size\"] == 64\n    assert PAPER_PROTOCOL[\"gan_image_size\"] == 512\n    assert PAPER_PROTOCOL[\"gan_kernel_size\"] == 4\n    assert PAPER_PROTOCOL[\"gan_stride\"] == 2\n    assert PAPER_PROTOCOL[\"critic_weight_clip\"] == 0.01\n    print(\"All explicitly reported numerical parameters pass the protocol assertions.\")\n\nassert_paper_protocol()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:13:31.112290Z","iopub.execute_input":"2026-07-10T17:13:31.112598Z","iopub.status.idle":"2026-07-10T17:13:31.118494Z","shell.execute_reply.started":"2026-07-10T17:13:31.112576Z","shell.execute_reply":"2026-07-10T17:13:31.117850Z"}},"outputs":[],"execution_count":null},{"id":"53e43149","cell_type":"markdown","source":"## 2. Locate the two datasets\n\nAttach these Kaggle inputs before running:\n\n1. `chest-xray-covid19-pneumonia`\n2. `rsna-pneumonia-detection-challenge`\n\nThe chest dataset supplies Normal, COVID-19, and Pneumonia. The RSNA dataset supplies Lung Opacity (`Target=1`).","metadata":{}},{"id":"80f6ccfd","cell_type":"code","source":"def directories_containing(start, token):\n    matches = []\n    for root, dirs, files in os.walk(start):\n        if token.lower() in root.lower():\n            matches.append(root)\n    return matches\n\ndef locate_cxr_root():\n    candidates = directories_containing(\"/kaggle/input\", \"chest-xray-covid19-pneumonia\")\n    for path in candidates:\n        path = Path(path)\n        if (path / \"Data\").is_dir():\n            return path / \"Data\"\n        if (path / \"train\").is_dir() and (path / \"test\").is_dir():\n            return path\n    return None\n\ndef locate_rsna_root():\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        if \"stage_2_train_labels.csv\" in files:\n            return Path(root)\n    return None\n\nCXR_ROOT = locate_cxr_root()\nRSNA_ROOT = locate_rsna_root()\n\nassert CXR_ROOT is not None, (\n    \"Chest X-ray dataset not found. Add the Kaggle dataset \"\n    \"'chest-xray-covid19-pneumonia'.\"\n)\nassert RSNA_ROOT is not None, (\n    \"RSNA dataset not found. Add 'rsna-pneumonia-detection-challenge' \"\n    \"and accept its competition rules.\"\n)\n\nRSNA_IMAGE_DIR = RSNA_ROOT / \"stage_2_train_images\"\nRSNA_LABEL_CSV = RSNA_ROOT / \"stage_2_train_labels.csv\"\n\nprint(\"CXR root :\", CXR_ROOT)\nprint(\"RSNA root:\", RSNA_ROOT)\n\nrsna_labels = pd.read_csv(RSNA_LABEL_CSV)\nrsna_unique = (\n    rsna_labels.groupby(\"patientId\", as_index=False)[\"Target\"]\n    .max()\n)\nprint(rsna_unique[\"Target\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:13:31.119617Z","iopub.execute_input":"2026-07-10T17:13:31.119842Z","iopub.status.idle":"2026-07-10T17:14:20.064142Z","shell.execute_reply.started":"2026-07-10T17:13:31.119822Z","shell.execute_reply":"2026-07-10T17:14:20.063300Z"}},"outputs":[],"execution_count":null},{"id":"7ae2fde1","cell_type":"markdown","source":"## 3. Load and randomly select exactly 25 real images per class\n\nImages from both existing chest-dataset folders are pooled before random selection, because the paper subsequently specifies a fresh stratified 80/20 random split.\n\nNo percentile windowing, contrast manipulation, rejection filtering, rotation, zoom, or other unreported augmentation is applied.","metadata":{}},{"id":"15522100","cell_type":"code","source":"GAN_IMAGE_SIZE = PAPER_PROTOCOL[\"gan_image_size\"]\n\ndef simple_dicom_to_uint8(ds):\n    arr = ds.pixel_array.astype(np.float32)\n    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n    arr = arr * slope + intercept\n\n    if str(getattr(ds, \"PhotometricInterpretation\", \"\")).upper() == \"MONOCHROME1\":\n        arr = arr.max() - arr\n\n    lo = float(np.min(arr))\n    hi = float(np.max(arr))\n    if hi <= lo:\n        return np.zeros_like(arr, dtype=np.uint8)\n    arr = (arr - lo) / (hi - lo)\n    return np.clip(arr * 255.0, 0, 255).astype(np.uint8)\n\ndef resize_grayscale_uint8(arr, size=GAN_IMAGE_SIZE):\n    arr = np.asarray(arr)\n    if arr.ndim == 3:\n        arr = arr[..., 0]\n    resized = tf.image.resize(\n        arr[..., None].astype(np.float32),\n        (size, size),\n        method=\"bilinear\",\n        antialias=True,\n    ).numpy()\n    return np.clip(resized, 0, 255).astype(np.uint8)\n\ndef load_chest_image(path):\n    image = keras.utils.load_img(path, color_mode=\"grayscale\")\n    arr = keras.utils.img_to_array(image)[..., 0]\n    return resize_grayscale_uint8(arr)\n\ndef load_rsna_image(patient_id):\n    path = RSNA_IMAGE_DIR / f\"{patient_id}.dcm\"\n    ds = pydicom.dcmread(path)\n    return resize_grayscale_uint8(simple_dicom_to_uint8(ds))\n\ndef chest_files(class_folder):\n    files = []\n    for split in (\"train\", \"test\"):\n        files.extend(glob.glob(str(CXR_ROOT / split / class_folder / \"*\")))\n    files = sorted([f for f in files if Path(f).is_file()])\n    return files\n\ndef choose_without_replacement(items, n, rng):\n    items = list(items)\n    if len(items) < n:\n        raise ValueError(f\"Requested {n} items, but only {len(items)} are available.\")\n    idx = rng.choice(len(items), size=n, replace=False)\n    return [items[i] for i in idx]\n\nrng = np.random.default_rng(SEED)\nn_real = PAPER_PROTOCOL[\"real_images_per_class\"]\n\nselected_sources = {}\nselected_sources[\"Normal\"] = choose_without_replacement(chest_files(\"NORMAL\"), n_real, rng)\nselected_sources[\"COVID-19\"] = choose_without_replacement(chest_files(\"COVID19\"), n_real, rng)\nselected_sources[\"Viral Pneumonia\"] = choose_without_replacement(chest_files(\"PNEUMONIA\"), n_real, rng)\n\nopacity_ids = rsna_unique.loc[rsna_unique[\"Target\"] == 1, \"patientId\"].astype(str).tolist()\nselected_sources[\"Lung Opacity\"] = choose_without_replacement(opacity_ids, n_real, rng)\n\nseed_images = {}\nfor class_name in CLASSES:\n    if class_name == \"Lung Opacity\":\n        seed_images[class_name] = np.stack(\n            [load_rsna_image(pid) for pid in selected_sources[class_name]]\n        )\n    else:\n        seed_images[class_name] = np.stack(\n            [load_chest_image(path) for path in selected_sources[class_name]]\n        )\n\nseed_manifest_rows = []\nfor class_name in CLASSES:\n    for source in selected_sources[class_name]:\n        seed_manifest_rows.append({\"class\": class_name, \"source\": str(source)})\nseed_manifest = pd.DataFrame(seed_manifest_rows)\nseed_manifest.to_csv(WORK_DIR / \"selected_real_images.csv\", index=False)\n\nfor class_name in CLASSES:\n    print(class_name, seed_images[class_name].shape)\n\nassert sum(len(seed_images[c]) for c in CLASSES) == 100\nassert all(len(seed_images[c]) == 25 for c in CLASSES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:27.524358Z","iopub.execute_input":"2026-07-10T17:15:27.525082Z","iopub.status.idle":"2026-07-10T17:15:36.292093Z","shell.execute_reply.started":"2026-07-10T17:15:27.525048Z","shell.execute_reply":"2026-07-10T17:15:36.291144Z"}},"outputs":[],"execution_count":null},{"id":"ecf95847","cell_type":"code","source":"fig, axes = plt.subplots(4, 5, figsize=(13, 10))\nfor row, class_name in enumerate(CLASSES):\n    for col in range(5):\n        axes[row, col].imshow(seed_images[class_name][col, ..., 0], cmap=\"gray\")\n        axes[row, col].axis(\"off\")\n        if col == 0:\n            axes[row, col].set_ylabel(class_name, rotation=0, labelpad=55)\nplt.suptitle(\"Randomly selected real seed images: 25 per class\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:36.293513Z","iopub.execute_input":"2026-07-10T17:15:36.293835Z","iopub.status.idle":"2026-07-10T17:15:37.462049Z","shell.execute_reply.started":"2026-07-10T17:15:36.293811Z","shell.execute_reply":"2026-07-10T17:15:37.461232Z"}},"outputs":[],"execution_count":null},{"id":"f92ca94d","cell_type":"markdown","source":"## 4. Classifier preprocessing and transfer-learning models\n\nThe paper describes feature-wise Z-normalization:\n\n\\[\n\\hat X_i = \f\nrac{X_i-\\mu_i}{\\sigma_i}.\n\\]\n\nThe implementation calculates a mean and standard-deviation image from the current training partition only and applies those statistics to training, validation, and testing images. Grayscale images are repeated into three channels for ImageNet-pretrained networks. No architecture-specific preprocessing is added because the paper reports Z-normalization instead.","metadata":{}},{"id":"6f56542e","cell_type":"code","source":"AUTOTUNE = tf.data.AUTOTUNE\n\ndef stratified_paper_split(X, y, origins):\n    # First preserve the reported 80/20 train/test split, then make validation\n    # from the training pool using the separately disclosed assumption.\n    X_train_pool, X_test, y_train_pool, y_test, o_train_pool, o_test = train_test_split(\n        X,\n        y,\n        origins,\n        train_size=PAPER_PROTOCOL[\"train_fraction\"],\n        test_size=PAPER_PROTOCOL[\"test_fraction\"],\n        random_state=SEED,\n        stratify=y,\n        shuffle=True,\n    )\n\n    val_fraction = UNREPORTED_ASSUMPTIONS[\"validation_fraction_within_training_pool\"]\n    X_train, X_val, y_train, y_val, o_train, o_val = train_test_split(\n        X_train_pool,\n        y_train_pool,\n        o_train_pool,\n        test_size=val_fraction,\n        random_state=SEED,\n        stratify=y_train_pool,\n        shuffle=True,\n    )\n    return {\n        \"X_train\": X_train,\n        \"y_train\": y_train,\n        \"origin_train\": o_train,\n        \"X_val\": X_val,\n        \"y_val\": y_val,\n        \"origin_val\": o_val,\n        \"X_test\": X_test,\n        \"y_test\": y_test,\n        \"origin_test\": o_test,\n    }\n\ndef compute_featurewise_z_stats(X_train, target_size, batch_size=16):\n    total = np.zeros((target_size, target_size, 1), dtype=np.float64)\n    total_sq = np.zeros_like(total)\n    count = 0\n\n    for start in range(0, len(X_train), batch_size):\n        batch = X_train[start:start + batch_size].astype(np.float32)\n        resized = tf.image.resize(\n            batch, (target_size, target_size),\n            method=\"bilinear\", antialias=True\n        ).numpy().astype(np.float64)\n        total += resized.sum(axis=0)\n        total_sq += np.square(resized).sum(axis=0)\n        count += len(resized)\n\n    mean = total / count\n    var = np.maximum(total_sq / count - np.square(mean), 0.0)\n    std = np.sqrt(var)\n    std[std < 1e-6] = 1.0\n    return mean.astype(np.float32), std.astype(np.float32)\n\ndef make_classifier_dataset(X, y, target_size, class_count, mean, std, training):\n    mean_t = tf.constant(mean, dtype=tf.float32)\n    std_t = tf.constant(std, dtype=tf.float32)\n\n    ds = tf.data.Dataset.from_tensor_slices((X, y))\n    if training:\n        ds = ds.shuffle(len(X), seed=SEED, reshuffle_each_iteration=True)\n\n    def prepare(image, label):\n        image = tf.cast(image, tf.float32)\n        image = tf.image.resize(\n            image, (target_size, target_size),\n            method=\"bilinear\", antialias=True\n        )\n        image = (image - mean_t) / std_t\n        image = tf.repeat(image, repeats=3, axis=-1)\n        label = tf.one_hot(label, depth=class_count)\n        return image, label\n\n    ds = ds.map(prepare, num_parallel_calls=AUTOTUNE)\n    ds = ds.batch(PAPER_PROTOCOL[\"classifier_batch_size\"], drop_remainder=False)\n    return ds.prefetch(AUTOTUNE)\n\ndef z_normalize_one(image, target_size, mean, std):\n    image = tf.image.resize(\n        tf.cast(image, tf.float32),\n        (target_size, target_size),\n        method=\"bilinear\", antialias=True,\n    )\n    image = (image - mean) / std\n    return tf.repeat(image, repeats=3, axis=-1).numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:37.463178Z","iopub.execute_input":"2026-07-10T17:15:37.463524Z","iopub.status.idle":"2026-07-10T17:15:37.477254Z","shell.execute_reply.started":"2026-07-10T17:15:37.463490Z","shell.execute_reply":"2026-07-10T17:15:37.476502Z"}},"outputs":[],"execution_count":null},{"id":"4e3cb3ce","cell_type":"code","source":"def load_imagenet_backbone(name, input_size):\n    try:\n        if name == \"VGG19\":\n            return keras.applications.VGG19(\n                include_top=False,\n                weights=\"imagenet\",\n                input_shape=(input_size, input_size, 3),\n            )\n        if name == \"InceptionV3\":\n            return keras.applications.InceptionV3(\n                include_top=False,\n                weights=\"imagenet\",\n                input_shape=(input_size, input_size, 3),\n            )\n        if name == \"Xception\":\n            return keras.applications.Xception(\n                include_top=False,\n                weights=\"imagenet\",\n                input_shape=(input_size, input_size, 3),\n            )\n    except Exception as exc:\n        raise RuntimeError(\n            \"ImageNet weights could not be loaded. Enable Kaggle Internet or \"\n            \"attach the corresponding Keras weight files. Falling back to \"\n            \"random weights would violate the paper.\"\n        ) from exc\n    raise ValueError(name)\n\ndef model_input_size(name):\n    return {\n        \"VGG19\": PAPER_PROTOCOL[\"vgg19_input\"],\n        \"InceptionV3\": PAPER_PROTOCOL[\"inceptionv3_input\"],\n        \"Xception\": PAPER_PROTOCOL[\"xception_input\"],\n    }[name]\n\ndef build_paper_classifier(name, class_count):\n    size = model_input_size(name)\n    backbone = load_imagenet_backbone(name, size)\n\n    # The paper freezes the pretrained VGG19 through conv5_pool and describes\n    # final-layer modification for the other pretrained backbones.\n    backbone.trainable = False\n    for layer in backbone.layers:\n        layer.trainable = False\n\n    inputs = layers.Input((size, size, 3), name=\"z_normalized_image\")\n    x = backbone(inputs, training=False)\n    x = layers.Activation(\"linear\", name=\"last_conv_features\")(x)\n\n    if name == \"VGG19\":\n        x = layers.AveragePooling2D(pool_size=(2, 2), name=\"average_pool\")(x)\n        x = layers.Flatten(name=\"flatten\")(x)\n    elif name == \"InceptionV3\":\n        x = layers.AveragePooling2D(pool_size=(2, 2), name=\"average_pool\")(x)\n        x = layers.Flatten(name=\"flatten\")(x)\n        x = layers.Dropout(\n            UNREPORTED_ASSUMPTIONS[\"inception_dropout_rate\"],\n            name=\"dropout\",\n        )(x)\n    elif name == \"Xception\":\n        p = PAPER_PROTOCOL[\"xception_pool_size\"]\n        x = layers.AveragePooling2D(pool_size=(p, p), name=\"average_pool_4x4\")(x)\n        x = layers.Flatten(name=\"flatten\")(x)\n    else:\n        raise ValueError(name)\n\n    outputs = layers.Dense(\n        class_count,\n        activation=\"softmax\",\n        name=\"dense_classification\",\n    )(x)\n\n    model = keras.Model(inputs, outputs, name=f\"{name}_paper\")\n    return model, backbone, size\n\ndef weighted_categorical_crossentropy(class_weights):\n    weights = tf.constant(np.asarray(class_weights, dtype=np.float32))\n\n    def loss(y_true, y_pred):\n        y_pred = tf.clip_by_value(\n            y_pred,\n            tf.keras.backend.epsilon(),\n            1.0 - tf.keras.backend.epsilon(),\n        )\n        return -tf.reduce_sum(\n            y_true * tf.math.log(y_pred) * weights,\n            axis=-1,\n        )\n    return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:37.478738Z","iopub.execute_input":"2026-07-10T17:15:37.479078Z","iopub.status.idle":"2026-07-10T17:15:37.496883Z","shell.execute_reply.started":"2026-07-10T17:15:37.479057Z","shell.execute_reply":"2026-07-10T17:15:37.496062Z"}},"outputs":[],"execution_count":null},{"id":"9f1e807e","cell_type":"markdown","source":"## 5. Experiment utilities\n\nThe paper describes six configurations: three binary and three multiclass models. It also states that the networks were initially trained on the 100 real images and retrained after adding the 1,600 GAN-generated images. The notebook therefore evaluates both the real-only and expanded stages.","metadata":{}},{"id":"ac7dfe3b","cell_type":"code","source":"def assemble_real_only_dataset():\n    X, y, origins = [], [], []\n    for class_id, class_name in enumerate(CLASSES):\n        X.extend(seed_images[class_name])\n        y.extend([class_id] * len(seed_images[class_name]))\n        origins.extend([\"real\"] * len(seed_images[class_name]))\n    return (\n        np.stack(X).astype(np.uint8),\n        np.asarray(y, dtype=np.int32),\n        np.asarray(origins),\n    )\n\ndef select_task(X, y, origins, task):\n    if task == \"binary\":\n        mask = np.isin(y, [CLASS_TO_ID[\"Normal\"], CLASS_TO_ID[\"COVID-19\"]])\n        X_task = X[mask]\n        y_task = y[mask].copy()  # already 0 and 1\n        origins_task = origins[mask]\n        class_names = [\"Normal\", \"COVID-19\"]\n    elif task == \"multiclass\":\n        X_task = X\n        y_task = y\n        origins_task = origins\n        class_names = CLASSES\n    else:\n        raise ValueError(task)\n    return X_task, y_task, origins_task, class_names\n\ndef experiment_tag(stage, task, model_name):\n    return f\"{stage}__{task}__{model_name.replace(' ', '_')}\"\n\nall_metric_rows = []\nexperiment_index = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:40.475180Z","iopub.execute_input":"2026-07-10T17:15:40.475633Z","iopub.status.idle":"2026-07-10T17:15:40.483191Z","shell.execute_reply.started":"2026-07-10T17:15:40.475604Z","shell.execute_reply":"2026-07-10T17:15:40.482469Z"}},"outputs":[],"execution_count":null},{"id":"21b7ec33","cell_type":"code","source":"def evaluate_predictions(stage, task, model_name, class_names, y_true, y_prob):\n    y_pred = np.argmax(y_prob, axis=1)\n    acc = accuracy_score(y_true, y_pred)\n    precision, recall, f1, _ = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=np.arange(len(class_names)),\n        zero_division=0,\n    )\n\n    rows = []\n    for class_id, class_name in enumerate(class_names):\n        binary_truth = (y_true == class_id).astype(np.int32)\n        fpr, tpr, _ = roc_curve(binary_truth, y_prob[:, class_id])\n        class_auc = auc(fpr, tpr)\n        rows.append({\n            \"Stage\": stage,\n            \"Task\": task,\n            \"Model\": model_name,\n            \"Class\": class_name,\n            \"Accuracy\": acc,\n            \"Precision_PPV\": precision[class_id],\n            \"Recall_Sensitivity\": recall[class_id],\n            \"F1\": f1[class_id],\n            \"AUC\": class_auc,\n        })\n    return rows, y_pred\n\ndef run_classifier_stage(stage, X, y, origins, tasks):\n    stage_records = []\n\n    for task in tasks:\n        X_task, y_task, origins_task, class_names = select_task(\n            X, y, origins, task\n        )\n        split = stratified_paper_split(X_task, y_task, origins_task)\n        class_count = len(class_names)\n\n        print(\"\\n\" + \"=\" * 80)\n        print(stage, task)\n        print(\"80/20 train pool/test sizes:\",\n              len(split[\"X_train\"]) + len(split[\"X_val\"]),\n              len(split[\"X_test\"]))\n        print(\"Final train/validation/test:\",\n              len(split[\"X_train\"]), len(split[\"X_val\"]), len(split[\"X_test\"]))\n        print(\"Test real fraction:\",\n              float(np.mean(split[\"origin_test\"] == \"real\")))\n\n        classes_for_weight = np.arange(class_count)\n        weights = compute_class_weight(\n            class_weight=\"balanced\",\n            classes=classes_for_weight,\n            y=split[\"y_train\"],\n        )\n        print(\"Class weights:\", dict(zip(class_names, weights)))\n\n        for model_name in RUN_CONTROL[\"models\"]:\n            print(\"\\n\", \"-\" * 25, model_name, \"-\" * 25)\n            tf.keras.backend.clear_session()\n            gc.collect()\n\n            model, backbone, input_size = build_paper_classifier(\n                model_name, class_count\n            )\n\n            mean, std = compute_featurewise_z_stats(\n                split[\"X_train\"], input_size\n            )\n\n            train_ds = make_classifier_dataset(\n                split[\"X_train\"], split[\"y_train\"], input_size,\n                class_count, mean, std, training=True\n            )\n            val_ds = make_classifier_dataset(\n                split[\"X_val\"], split[\"y_val\"], input_size,\n                class_count, mean, std, training=False\n            )\n            test_ds = make_classifier_dataset(\n                split[\"X_test\"], split[\"y_test\"], input_size,\n                class_count, mean, std, training=False\n            )\n\n            # No learning-rate schedule, label smoothing, fine-tuning stage,\n            # or early stopping is added because the paper does not report them.\n            model.compile(\n                optimizer=keras.optimizers.Adam(),  # Keras defaults\n                loss=weighted_categorical_crossentropy(weights),\n                metrics=[\"accuracy\"],\n            )\n\n            tag = experiment_tag(stage, task, model_name)\n            history_csv = WORK_DIR / \"histories\" / f\"{tag}.csv\"\n            callbacks = [\n                keras.callbacks.CSVLogger(history_csv),\n                keras.callbacks.TerminateOnNaN(),\n            ]\n\n            history = model.fit(\n                train_ds,\n                validation_data=val_ds,\n                epochs=CLASSIFIER_EPOCHS,\n                callbacks=callbacks,\n                verbose=2,\n            )\n\n            y_prob = model.predict(test_ds, verbose=1)\n            metric_rows, y_pred = evaluate_predictions(\n                stage, task, model_name, class_names,\n                split[\"y_test\"], y_prob\n            )\n            all_metric_rows.extend(metric_rows)\n\n            pred_path = WORK_DIR / \"predictions\" / f\"{tag}.npz\"\n            np.savez_compressed(\n                pred_path,\n                y_true=split[\"y_test\"],\n                y_prob=y_prob,\n                y_pred=y_pred,\n                X_test=split[\"X_test\"],\n                origin_test=split[\"origin_test\"],\n                class_names=np.asarray(class_names),\n            )\n\n            stats_path = WORK_DIR / \"predictions\" / f\"{tag}_zstats.npz\"\n            np.savez_compressed(stats_path, mean=mean, std=std)\n\n            model_path = None\n            if stage == \"augmented\" and task == \"multiclass\":\n                # Save weights rather than a compiled model because the paper's\n                # weighted categorical loss is implemented as a local closure.\n                # This avoids Keras serialization errors while preserving every\n                # trained parameter exactly.\n                model_path = WORK_DIR / \"models\" / f\"{tag}.weights.h5\"\n                model.save_weights(model_path)\n\n            stage_records.append({\n                \"Stage\": stage,\n                \"Task\": task,\n                \"Model\": model_name,\n                \"InputSize\": input_size,\n                \"Epochs\": CLASSIFIER_EPOCHS,\n                \"BatchSize\": PAPER_PROTOCOL[\"classifier_batch_size\"],\n                \"Optimizer\": \"Adam defaults\",\n                \"HistoryCSV\": str(history_csv),\n                \"PredictionFile\": str(pred_path),\n                \"ZStatsFile\": str(stats_path),\n                \"ModelFile\": str(model_path) if model_path else \"\",\n            })\n\n            print(classification_report(\n                split[\"y_test\"], y_pred,\n                target_names=class_names,\n                zero_division=0,\n            ))\n\n            del model, backbone, train_ds, val_ds, test_ds, y_prob\n            gc.collect()\n\n    return stage_records","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:15:54.458292Z","iopub.execute_input":"2026-07-10T17:15:54.459142Z","iopub.status.idle":"2026-07-10T17:15:54.475213Z","shell.execute_reply.started":"2026-07-10T17:15:54.459106Z","shell.execute_reply":"2026-07-10T17:15:54.474234Z"}},"outputs":[],"execution_count":null},{"id":"3982fa77","cell_type":"markdown","source":"## 6. Initial transfer-learning stage: 100 real images","metadata":{}},{"id":"b0aca7e6","cell_type":"code","source":"real_X, real_y, real_origins = assemble_real_only_dataset()\nprint(real_X.shape, np.bincount(real_y), np.unique(real_origins, return_counts=True))\n\ntasks_to_run = []\nif RUN_CONTROL[\"run_binary_models\"]:\n    tasks_to_run.append(\"binary\")\nif RUN_CONTROL[\"run_multiclass_models\"]:\n    tasks_to_run.append(\"multiclass\")\n\nif RUN_CONTROL[\"run_real_only_baseline\"]:\n    baseline_records = run_classifier_stage(\n        \"real_only\", real_X, real_y, real_origins, tasks_to_run\n    )\n    experiment_index.extend(baseline_records)\nelse:\n    print(\"Real-only baseline skipped by RUN_CONTROL.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T17:16:01.155841Z","iopub.execute_input":"2026-07-10T17:16:01.156367Z"}},"outputs":[],"execution_count":null},{"id":"d592991c","cell_type":"markdown","source":"## 7. Paper OT-WGAN\n\n### Generator\n\n- Random noise input\n- Dense projection and reshape to **64×64**\n- Three `Conv2DTranspose` blocks:\n  - 64→128\n  - 128→256\n  - 256→512\n- Kernel **4**, stride **2**\n- LeakyReLU\n- Final single-channel image\n\n### Critic\n\n- 512→256→128→64 with stride-2 convolution\n- LeakyReLU\n- Linear scalar output\n- Wasserstein objective\n- Weight clipping to **[-0.01, 0.01] after every critic update**\n\n### Sinkhorn term\n\nThe implementation follows the paper's displayed entropy-regularized OT objective rather than the debiased Sinkhorn divergence used in the previous notebook.","metadata":{}},{"id":"e0fdd5a3","cell_type":"code","source":"GAN_SIZE = PAPER_PROTOCOL[\"gan_image_size\"]\nLATENT_DIM = UNREPORTED_ASSUMPTIONS[\"gan_latent_dim\"]\nGAN_KERNEL = PAPER_PROTOCOL[\"gan_kernel_size\"]\nGAN_STRIDE = PAPER_PROTOCOL[\"gan_stride\"]\n\ndef build_paper_generator():\n    start = PAPER_PROTOCOL[\"generator_start_size\"]\n    c0 = UNREPORTED_ASSUMPTIONS[\"generator_start_channels\"]\n\n    z = layers.Input((LATENT_DIM,), name=\"random_noise\")\n    x = layers.Dense(start * start * c0, name=\"dense_projection\")(z)\n    x = layers.Reshape((start, start, c0), name=\"reshape_64x64\")(x)\n    x = layers.LeakyReLU(negative_slope=0.2)(x)\n\n    for filters, output_size in [(32, 128), (16, 256), (8, 512)]:\n        x = layers.Conv2DTranspose(\n            filters,\n            kernel_size=GAN_KERNEL,\n            strides=GAN_STRIDE,\n            padding=\"same\",\n            name=f\"transpose_conv_to_{output_size}\",\n        )(x)\n        x = layers.LeakyReLU(\n            negative_slope=0.2,\n            name=f\"leaky_relu_{output_size}\",\n        )(x)\n\n    output = layers.Conv2DTranspose(\n        1,\n        kernel_size=GAN_KERNEL,\n        strides=1,\n        padding=\"same\",\n        activation=\"linear\",\n        name=\"generated_z_normalized_image\",\n    )(x)\n\n    return keras.Model(z, output, name=\"paper_generator\")\n\ndef build_paper_critic():\n    image = layers.Input((GAN_SIZE, GAN_SIZE, 1), name=\"image_512\")\n    x = image\n    for filters, output_size in [(8, 256), (16, 128), (32, 64)]:\n        x = layers.Conv2D(\n            filters,\n            kernel_size=GAN_KERNEL,\n            strides=GAN_STRIDE,\n            padding=\"same\",\n            name=f\"downsample_to_{output_size}\",\n        )(x)\n        x = layers.LeakyReLU(\n            negative_slope=0.2,\n            name=f\"critic_leaky_relu_{output_size}\",\n        )(x)\n\n    x = layers.Flatten(name=\"critic_flatten\")(x)\n    score = layers.Dense(1, activation=\"linear\", name=\"linear_critic_output\")(x)\n    return keras.Model(image, score, name=\"paper_critic\")\n\n_g = build_paper_generator()\n_c = build_paper_critic()\nprint(\"Generator:\", _g.input_shape, \"->\", _g.output_shape)\nprint(\"Critic   :\", _c.input_shape, \"->\", _c.output_shape)\n\nassert _g.output_shape[1:3] == (512, 512)\nassert _c.input_shape[1:3] == (512, 512)\nassert tuple(_c.get_layer(\"downsample_to_64\").output.shape[1:3]) == (64, 64)\n\nfor layer in _g.layers:\n    if isinstance(layer, layers.Conv2DTranspose) and layer.name.startswith(\"transpose_conv_to\"):\n        assert layer.kernel_size == (4, 4)\n        assert layer.strides == (2, 2)\n\ndel _g, _c\ngc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"1996e67c","cell_type":"code","source":"def entropy_regularized_sinkhorn_cost(real, fake):\n    # Equation (5)-style regularized OT cost over the two empirical batches.\n    x = tf.reshape(real, (tf.shape(real)[0], -1))\n    y = tf.reshape(fake, (tf.shape(fake)[0], -1))\n\n    # Mean squared Euclidean transport cost. Division by feature count keeps\n    # the numerical scale independent of the 512x512 dimensionality.\n    feature_count = tf.cast(tf.shape(x)[1], tf.float32)\n    x2 = tf.reduce_sum(tf.square(x), axis=1, keepdims=True)\n    y2 = tf.reduce_sum(tf.square(y), axis=1, keepdims=True)\n    cost = tf.maximum(\n        x2 + tf.transpose(y2) - 2.0 * tf.matmul(x, y, transpose_b=True),\n        0.0,\n    ) / feature_count\n\n    n = tf.shape(x)[0]\n    m = tf.shape(y)[0]\n    a = tf.ones((n, 1), dtype=tf.float32) / tf.cast(n, tf.float32)\n    b = tf.ones((m, 1), dtype=tf.float32) / tf.cast(m, tf.float32)\n\n    epsilon = tf.constant(\n        UNREPORTED_ASSUMPTIONS[\"sinkhorn_epsilon\"],\n        dtype=tf.float32,\n    )\n    kernel = tf.exp(-cost / epsilon) + 1e-9\n    u = tf.ones_like(a)\n    v = tf.ones_like(b)\n\n    for _ in range(UNREPORTED_ASSUMPTIONS[\"sinkhorn_iterations\"]):\n        u = a / (tf.matmul(kernel, v) + 1e-8)\n        v = b / (tf.matmul(kernel, u, transpose_a=True) + 1e-8)\n\n    transport = u * kernel * tf.transpose(v)\n    entropy = tf.reduce_sum(\n        transport * (tf.math.log(transport + 1e-9) - 1.0)\n    )\n    return tf.reduce_sum(transport * cost) + epsilon * entropy","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"46e16380","cell_type":"code","source":"class PaperOTWGANTrainer:\n    def __init__(self):\n        self.generator = build_paper_generator()\n        self.critic = build_paper_critic()\n\n        # The paper reports Adam with default parameters and gives no separate\n        # GAN optimizer settings; defaults are used here as the literal choice.\n        self.generator_optimizer = keras.optimizers.Adam()\n        self.critic_optimizer = keras.optimizers.Adam()\n\n    @tf.function(reduce_retracing=True)\n    def critic_step(self, real):\n        batch_size = tf.shape(real)[0]\n        z = tf.random.normal(tf.stack([batch_size, LATENT_DIM]))\n\n        with tf.GradientTape() as tape:\n            fake = self.generator(z, training=True)\n            fake = tf.stop_gradient(fake)\n\n            real_score = self.critic(real, training=True)\n            fake_score = self.critic(fake, training=True)\n\n            # Minimize fake - real == maximize real - fake.\n            critic_loss = (\n                tf.reduce_mean(fake_score) - tf.reduce_mean(real_score)\n            )\n\n        gradients = tape.gradient(\n            critic_loss, self.critic.trainable_variables\n        )\n        self.critic_optimizer.apply_gradients(\n            zip(gradients, self.critic.trainable_variables)\n        )\n\n        # Paper-required clipping after every critic batch.\n        clip_value = PAPER_PROTOCOL[\"critic_weight_clip\"]\n        for variable in self.critic.trainable_variables:\n            variable.assign(\n                tf.clip_by_value(variable, -clip_value, clip_value)\n            )\n\n        wasserstein_estimate = (\n            tf.reduce_mean(real_score) - tf.reduce_mean(fake_score)\n        )\n        return critic_loss, wasserstein_estimate\n\n    @tf.function(reduce_retracing=True)\n    def generator_step(self, real):\n        batch_size = tf.shape(real)[0]\n        z = tf.random.normal(tf.stack([batch_size, LATENT_DIM]))\n\n        with tf.GradientTape() as tape:\n            fake = self.generator(z, training=True)\n            fake_score = self.critic(fake, training=False)\n\n            wasserstein_generator_loss = -tf.reduce_mean(fake_score)\n            sinkhorn_loss = entropy_regularized_sinkhorn_cost(real, fake)\n\n            generator_loss = (\n                wasserstein_generator_loss\n                + UNREPORTED_ASSUMPTIONS[\"sinkhorn_weight\"] * sinkhorn_loss\n            )\n\n        gradients = tape.gradient(\n            generator_loss, self.generator.trainable_variables\n        )\n        self.generator_optimizer.apply_gradients(\n            zip(gradients, self.generator.trainable_variables)\n        )\n        return generator_loss, wasserstein_generator_loss, sinkhorn_loss\n\n    def fit(self, class_name, real_uint8):\n        real = real_uint8.astype(np.float32)\n\n        # Feature-wise Z-normalization over the 25 real images, as reported.\n        mean = real.mean(axis=0, keepdims=True)\n        std = real.std(axis=0, keepdims=True)\n        std[std < 1e-6] = 1.0\n        real_z = (real - mean) / std\n\n        batch_size = UNREPORTED_ASSUMPTIONS[\"gan_batch_size\"]\n        critic_updates = UNREPORTED_ASSUMPTIONS[\n            \"critic_updates_per_generator\"\n        ]\n\n        history_rows = []\n        for epoch in range(GAN_EPOCHS):\n            for _ in range(critic_updates):\n                idx = np.random.choice(\n                    len(real_z),\n                    size=batch_size,\n                    replace=True,\n                )\n                critic_loss, w_est = self.critic_step(\n                    tf.convert_to_tensor(real_z[idx], dtype=tf.float32)\n                )\n\n            idx = np.random.choice(\n                len(real_z),\n                size=batch_size,\n                replace=True,\n            )\n            g_loss, g_w_loss, sinkhorn_loss = self.generator_step(\n                tf.convert_to_tensor(real_z[idx], dtype=tf.float32)\n            )\n\n            if epoch % 25 == 0 or epoch == GAN_EPOCHS - 1:\n                row = {\n                    \"epoch\": epoch + 1,\n                    \"critic_loss\": float(critic_loss),\n                    \"wasserstein_estimate\": float(w_est),\n                    \"generator_loss\": float(g_loss),\n                    \"generator_wasserstein_loss\": float(g_w_loss),\n                    \"sinkhorn_loss\": float(sinkhorn_loss),\n                }\n                history_rows.append(row)\n                print(\n                    f\"[{class_name}] epoch {epoch + 1:4d}/{GAN_EPOCHS} \"\n                    f\"C={row['critic_loss']:.5f} \"\n                    f\"G={row['generator_loss']:.5f} \"\n                    f\"OT={row['sinkhorn_loss']:.5f}\"\n                )\n\n        return (\n            pd.DataFrame(history_rows),\n            mean.astype(np.float32),\n            std.astype(np.float32),\n        )\n\n    def generate(self, n, mean, std):\n        batches = []\n        for start in range(0, n, 16):\n            current = min(16, n - start)\n            z = tf.random.normal((current, LATENT_DIM))\n            fake_z = self.generator(z, training=False).numpy()\n            fake = fake_z * std + mean\n            batches.append(np.clip(fake, 0, 255).astype(np.uint8))\n        return np.concatenate(batches, axis=0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"c883bc6c","cell_type":"markdown","source":"## 8. Train one class-specific OT-WGAN per class and generate exactly 400 images per class","metadata":{}},{"id":"cee6b824","cell_type":"code","source":"synthetic_images = {}\ngan_history_files = {}\n\nif RUN_CONTROL[\"run_gan_training\"]:\n    for class_name in CLASSES:\n        safe_name = class_name.lower().replace(\" \", \"_\").replace(\"-\", \"\")\n        cache_path = WORK_DIR / \"gan_cache\" / f\"{safe_name}_synthetic.npz\"\n        history_path = WORK_DIR / \"gan_cache\" / f\"{safe_name}_history.csv\"\n\n        if cache_path.exists():\n            cached = np.load(cache_path)\n            synthetic_images[class_name] = cached[\"images\"]\n            print(class_name, \"loaded from cache:\", synthetic_images[class_name].shape)\n            gan_history_files[class_name] = str(history_path)\n            continue\n\n        print(\"\\n\" + \"=\" * 80)\n        print(\"Training class-specific OT-WGAN:\", class_name)\n\n        tf.keras.backend.clear_session()\n        gc.collect()\n\n        trainer = PaperOTWGANTrainer()\n        history, gan_mean, gan_std = trainer.fit(\n            class_name, seed_images[class_name]\n        )\n        generated = trainer.generate(\n            PAPER_PROTOCOL[\"generated_images_per_class\"],\n            gan_mean,\n            gan_std,\n        )\n\n        assert generated.shape == (\n            PAPER_PROTOCOL[\"generated_images_per_class\"],\n            GAN_SIZE,\n            GAN_SIZE,\n            1,\n        )\n\n        np.savez_compressed(cache_path, images=generated)\n        history.to_csv(history_path, index=False)\n\n        synthetic_images[class_name] = generated\n        gan_history_files[class_name] = str(history_path)\n\n        del trainer, generated, history, gan_mean, gan_std\n        gc.collect()\nelse:\n    for class_name in CLASSES:\n        safe_name = class_name.lower().replace(\" \", \"_\").replace(\"-\", \"\")\n        cache_path = WORK_DIR / \"gan_cache\" / f\"{safe_name}_synthetic.npz\"\n        if not cache_path.exists():\n            raise FileNotFoundError(\n                f\"GAN training is disabled, but cache is missing: {cache_path}\"\n            )\n        synthetic_images[class_name] = np.load(cache_path)[\"images\"]\n\ntotal_generated = sum(len(synthetic_images[c]) for c in CLASSES)\nprint(\"Total generated:\", total_generated)\nassert total_generated == PAPER_PROTOCOL[\"generated_images_total\"]\nassert all(len(synthetic_images[c]) == 400 for c in CLASSES)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"bfe2645e","cell_type":"code","source":"fig, axes = plt.subplots(4, 6, figsize=(15, 10))\nfor row, class_name in enumerate(CLASSES):\n    axes[row, 0].imshow(seed_images[class_name][0, ..., 0], cmap=\"gray\")\n    axes[row, 0].set_title(\"Real\")\n    axes[row, 0].axis(\"off\")\n    axes[row, 0].set_ylabel(class_name, rotation=0, labelpad=55)\n\n    for col in range(1, 6):\n        axes[row, col].imshow(\n            synthetic_images[class_name][col - 1, ..., 0],\n            cmap=\"gray\",\n        )\n        axes[row, col].set_title(\"Generated\")\n        axes[row, col].axis(\"off\")\n\nplt.suptitle(\"Real versus OT-WGAN generated images\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"31f8c8e4","cell_type":"markdown","source":"## 9. Expanded dataset and transfer-model retraining","metadata":{}},{"id":"d70327b2","cell_type":"code","source":"def assemble_augmented_dataset():\n    X, y, origins = [], [], []\n    for class_id, class_name in enumerate(CLASSES):\n        X.extend(seed_images[class_name])\n        y.extend([class_id] * len(seed_images[class_name]))\n        origins.extend([\"real\"] * len(seed_images[class_name]))\n\n        X.extend(synthetic_images[class_name])\n        y.extend([class_id] * len(synthetic_images[class_name]))\n        origins.extend([\"synthetic\"] * len(synthetic_images[class_name]))\n\n    return (\n        np.stack(X).astype(np.uint8),\n        np.asarray(y, dtype=np.int32),\n        np.asarray(origins),\n    )\n\naugmented_X, augmented_y, augmented_origins = assemble_augmented_dataset()\n\nprint(\"Expanded dataset shape:\", augmented_X.shape)\nprint(\"Class counts:\", np.bincount(augmented_y))\nprint(\"Origin counts:\", dict(zip(\n    *np.unique(augmented_origins, return_counts=True)\n)))\n\nassert len(augmented_X) == 1700\nassert np.all(np.bincount(augmented_y) == 425)\n\nif RUN_CONTROL[\"run_augmented_models\"]:\n    augmented_records = run_classifier_stage(\n        \"augmented\",\n        augmented_X,\n        augmented_y,\n        augmented_origins,\n        tasks_to_run,\n    )\n    experiment_index.extend(augmented_records)\nelse:\n    print(\"Augmented model training skipped by RUN_CONTROL.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"51651bda","cell_type":"markdown","source":"## 10. Save and display the complete result table","metadata":{}},{"id":"7d07fccd","cell_type":"code","source":"metrics_df = pd.DataFrame(all_metric_rows)\nindex_df = pd.DataFrame(experiment_index)\n\nmetrics_path = WORK_DIR / \"paper_protocol_metrics.csv\"\nindex_path = WORK_DIR / \"experiment_index.csv\"\n\nmetrics_df.to_csv(metrics_path, index=False)\nindex_df.to_csv(index_path, index=False)\n\ndisplay(metrics_df.round(4))\ndisplay(index_df)\n\nmodel_accuracy = (\n    metrics_df.groupby([\"Stage\", \"Task\", \"Model\"], sort=False)[\"Accuracy\"]\n    .first()\n    .reset_index()\n)\nprint(\"\\nModel-level accuracy\")\ndisplay(model_accuracy.round(4))\n\nprint(\"Saved:\", metrics_path)\nprint(\"Saved:\", index_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"93e2eda7","cell_type":"markdown","source":"## 11. Training and validation curves","metadata":{}},{"id":"e30b80e8","cell_type":"code","source":"if len(index_df):\n    for stage in index_df[\"Stage\"].unique():\n        for task in index_df[\"Task\"].unique():\n            subset = index_df[\n                (index_df[\"Stage\"] == stage) &\n                (index_df[\"Task\"] == task)\n            ]\n            if subset.empty:\n                continue\n\n            fig, axes = plt.subplots(1, len(subset), figsize=(6 * len(subset), 4))\n            if len(subset) == 1:\n                axes = [axes]\n\n            for ax, (_, row) in zip(axes, subset.iterrows()):\n                history = pd.read_csv(row[\"HistoryCSV\"])\n                ax.plot(history[\"epoch\"], history[\"accuracy\"], label=\"Train Accuracy\")\n                ax.plot(history[\"epoch\"], history[\"val_accuracy\"], label=\"Validation Accuracy\")\n                ax.set_title(f\"{stage} | {task} | {row['Model']}\")\n                ax.set_xlabel(\"Epoch\")\n                ax.set_ylabel(\"Accuracy\")\n                ax.legend()\n\n            plt.tight_layout()\n            plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"8c1b855a","cell_type":"markdown","source":"## 12. Confusion matrices and ROC curves","metadata":{}},{"id":"5b0a90ea","cell_type":"code","source":"for _, row in index_df.iterrows():\n    pred = np.load(row[\"PredictionFile\"], allow_pickle=True)\n    y_true = pred[\"y_true\"]\n    y_prob = pred[\"y_prob\"]\n    y_pred = pred[\"y_pred\"]\n    class_names = pred[\"class_names\"].tolist()\n\n    cm = confusion_matrix(y_true, y_pred, normalize=\"true\")\n\n    plt.figure(figsize=(6, 5))\n    plt.imshow(cm, interpolation=\"nearest\")\n    plt.title(f\"{row['Stage']} | {row['Task']} | {row['Model']}\")\n    plt.colorbar()\n    ticks = np.arange(len(class_names))\n    plt.xticks(ticks, class_names, rotation=35, ha=\"right\")\n    plt.yticks(ticks, class_names)\n    plt.xlabel(\"Predicted label\")\n    plt.ylabel(\"True label\")\n\n    threshold = cm.max() / 2.0 if cm.size else 0.5\n    for i in range(cm.shape[0]):\n        for j in range(cm.shape[1]):\n            plt.text(\n                j, i, f\"{cm[i, j]:.3f}\",\n                ha=\"center\", va=\"center\",\n            )\n    plt.tight_layout()\n    plt.show()\n\n    plt.figure(figsize=(6, 5))\n    for class_id, class_name in enumerate(class_names):\n        truth = (y_true == class_id).astype(np.int32)\n        fpr, tpr, _ = roc_curve(truth, y_prob[:, class_id])\n        plt.plot(fpr, tpr, label=f\"{class_name} (AUC={auc(fpr, tpr):.3f})\")\n    plt.plot([0, 1], [0, 1], linestyle=\"--\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.title(f\"ROC | {row['Stage']} | {row['Task']} | {row['Model']}\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"1a55539d","cell_type":"markdown","source":"## 13. Grad-CAM for the three augmented multiclass models\n\nThe heatmap is generated from `last_conv_features`, which is the final convolutional output of each pretrained backbone. One test image from each class is visualized for every model, both as a heatmap and as an overlay.","metadata":{}},{"id":"c1eacf77","cell_type":"code","source":"def gradcam_heatmap(model, preprocessed_image):\n    grad_model = keras.Model(\n        model.inputs,\n        [\n            model.get_layer(\"last_conv_features\").output,\n            model.output,\n        ],\n    )\n\n    image_tensor = tf.convert_to_tensor(\n        preprocessed_image[None, ...],\n        dtype=tf.float32,\n    )\n\n    with tf.GradientTape() as tape:\n        conv_output, predictions = grad_model(image_tensor, training=False)\n        class_id = tf.argmax(predictions[0])\n        class_score = predictions[:, class_id]\n\n    gradients = tape.gradient(class_score, conv_output)\n    weights = tf.reduce_mean(gradients, axis=(1, 2), keepdims=True)\n    heatmap = tf.reduce_sum(weights * conv_output, axis=-1)[0]\n    heatmap = tf.maximum(heatmap, 0)\n    heatmap = heatmap / (tf.reduce_max(heatmap) + 1e-8)\n    return heatmap.numpy(), int(class_id.numpy())\n\ngradcam_rows = index_df[\n    (index_df[\"Stage\"] == \"augmented\") &\n    (index_df[\"Task\"] == \"multiclass\") &\n    (index_df[\"ModelFile\"] != \"\")\n]\n\nfor _, row in gradcam_rows.iterrows():\n    tf.keras.backend.clear_session()\n    gc.collect()\n\n    pred = np.load(row[\"PredictionFile\"], allow_pickle=True)\n    stats = np.load(row[\"ZStatsFile\"])\n    X_test = pred[\"X_test\"]\n    y_true = pred[\"y_true\"]\n    class_names = pred[\"class_names\"].tolist()\n\n    model, _, _ = build_paper_classifier(\n        row[\"Model\"], len(class_names)\n    )\n    model.load_weights(row[\"ModelFile\"])\n    mean = stats[\"mean\"]\n    std = stats[\"std\"]\n    input_size = int(row[\"InputSize\"])\n\n    fig_heat, axes_heat = plt.subplots(1, len(class_names), figsize=(14, 4))\n    fig_overlay, axes_overlay = plt.subplots(1, len(class_names), figsize=(14, 4))\n\n    for class_id, class_name in enumerate(class_names):\n        idx = np.where(y_true == class_id)[0][0]\n        prepared = z_normalize_one(\n            X_test[idx], input_size, mean, std\n        )\n        heatmap, predicted_class = gradcam_heatmap(model, prepared)\n        heatmap_resized = tf.image.resize(\n            heatmap[..., None], (input_size, input_size)\n        ).numpy().squeeze()\n\n        original_resized = tf.image.resize(\n            X_test[idx].astype(np.float32),\n            (input_size, input_size),\n            method=\"bilinear\",\n            antialias=True,\n        ).numpy().squeeze()\n\n        axes_heat[class_id].imshow(heatmap_resized)\n        axes_heat[class_id].set_title(class_name)\n        axes_heat[class_id].axis(\"off\")\n\n        axes_overlay[class_id].imshow(original_resized, cmap=\"gray\")\n        axes_overlay[class_id].imshow(heatmap_resized, alpha=0.42)\n        axes_overlay[class_id].set_title(\n            f\"{class_name}\\nPred: {class_names[predicted_class]}\"\n        )\n        axes_overlay[class_id].axis(\"off\")\n\n    fig_heat.suptitle(f\"{row['Model']} — last-layer Grad-CAM heatmaps\")\n    fig_overlay.suptitle(f\"{row['Model']} — Grad-CAM overlays\")\n    fig_heat.tight_layout()\n    fig_overlay.tight_layout()\n    plt.show()\n\n    del model\n    gc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"79f113fd","cell_type":"markdown","source":"## 14. Paper-result reference\n\nThe paper reports the following four-class accuracies after augmentation:\n\n- VGG19: **96.94%**\n- InceptionV3: **95.35%**\n- Xception: **92.62%**\n\nThe notebook does not hard-code, fabricate, round, or overwrite measured outputs to force those values. Reaching the same numbers still depends on the authors' unreported GAN and validation settings, their exact randomly chosen images, and their precise implementation.","metadata":{}}]}