{"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":"89ba17a1","cell_type":"markdown","source":"# Leakage-free reproduction of the OT-GAN chest X-ray paper\n\nThis notebook implements the four-class experiment from **“Optimal Transport Theory-based GAN for Medical Image Augmentation and Classification”** while correcting the evaluation leakage in the original workflow.\n\n## Corrections made\n\n1. The 25 real images selected for each class are split **before any GAN or classifier is trained**.\n2. Per class, the clean split is:\n   - 16 real training images\n   - 4 real validation images\n   - 5 real test images\n3. Each class-specific OT-WGAN is trained using only the 16 real training images from that class.\n4. All 1,600 generated images are used only in the classifier training set.\n5. Validation and test sets contain only unseen real images.\n6. The real-only and real-plus-synthetic classifiers use the exact same real validation and test sets.\n7. The Sinkhorn term is implemented as a **debiased Sinkhorn divergence**:\n\n\\[\nS_\\varepsilon(\\mu,\\nu)=OT_\\varepsilon(\\mu,\\nu)\n-\\frac{1}{2}OT_\\varepsilon(\\mu,\\mu)\n-\\frac{1}{2}OT_\\varepsilon(\\nu,\\nu).\n\\]\n\n## Important limitation\n\nThe clean test set contains 20 real images in total, so overall accuracy can only change in steps of 5 percentage points. Therefore, the paper's exact reported value of 96.94% cannot be reproduced from this clean 25-images-per-class split. This notebook evaluates whether synthetic augmentation helps on unseen real images; it does not force the paper's reported numbers.","metadata":{}},{"id":"11c51bd5","cell_type":"code","source":"import os\nimport gc\nimport glob\nimport json\nimport random\nimport hashlib\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\nfrom IPython.display import display\n\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-14T15:51:23.440470Z","iopub.execute_input":"2026-07-14T15:51:23.441324Z","iopub.status.idle":"2026-07-14T15:51:23.448500Z","shell.execute_reply.started":"2026-07-14T15:51:23.441293Z","shell.execute_reply":"2026-07-14T15:51:23.447621Z"}},"outputs":[],"execution_count":null},{"id":"98901822","cell_type":"markdown","source":"## 1. Configuration\n\nValues under `PAPER_PROTOCOL` are stated by the paper. Values under `IMPLEMENTATION_ASSUMPTIONS` are required because the paper does not report them.","metadata":{}},{"id":"d4aabaf2","cell_type":"code","source":"PAPER_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    \"reported_train_fraction\": 0.80,\n    \"reported_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}\n\nIMPLEMENTATION_ASSUMPTIONS = {\n    \"random_seed\": 42,\n    # The 20-image development portion is separated into 16 train + 4 validation\n    # before the GAN sees any image.\n    \"real_train_per_class\": 16,\n    \"real_validation_per_class\": 4,\n    \"real_test_per_class\": 5,\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    \"sinkhorn_pool_size\": 32,\n    \"generator_start_channels\": 32,\n    \"inception_dropout_rate\": 0.50,\n    \"gradient_clip_norm\": 1.0,\n    \"gan_log_every\": 10,\n    \"use_best_validation_checkpoint\": True,\n}\n\nRUN_CONTROL = {\n    \"quick_debug\": False,\n    \"run_real_only_baseline\": True,\n    \"run_gan_training\": True,\n    \"run_augmented_models\": True,\n    \"models\": [\"VGG19\", \"InceptionV3\", \"Xception\"],\n}\n\nif RUN_CONTROL[\"quick_debug\"]:\n    print(\"WARNING: quick-debug mode is active; this is not the full protocol.\")\n    CLASSIFIER_EPOCHS = 2\n    GAN_EPOCHS = 2\nelse:\n    CLASSIFIER_EPOCHS = PAPER_PROTOCOL[\"classifier_epochs\"]\n    GAN_EPOCHS = IMPLEMENTATION_ASSUMPTIONS[\"gan_epochs\"]\n\nSEED = IMPLEMENTATION_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_corrected_reproduction\")\nfor subdir in [\"models\", \"histories\", \"predictions\", \"gan_cache\"]:\n    (WORK_DIR / subdir).mkdir(parents=True, exist_ok=True)\n\ndisplay(pd.Series(PAPER_PROTOCOL, name=\"paper-reported value\").to_frame())\ndisplay(pd.Series(IMPLEMENTATION_ASSUMPTIONS, name=\"implementation assumption\").to_frame())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:23.450918Z","iopub.execute_input":"2026-07-14T15:51:23.451290Z","iopub.status.idle":"2026-07-14T15:51:23.476182Z","shell.execute_reply.started":"2026-07-14T15:51:23.451263Z","shell.execute_reply":"2026-07-14T15:51:23.475136Z"}},"outputs":[],"execution_count":null},{"id":"bb01d3a1","cell_type":"code","source":"def assert_protocol_configuration():\n    assert PAPER_PROTOCOL[\"real_images_per_class\"] == 25\n    assert PAPER_PROTOCOL[\"generated_images_per_class\"] * NUM_CLASSES == 1600\n    assert PAPER_PROTOCOL[\"classifier_batch_size\"] == 16\n    assert PAPER_PROTOCOL[\"classifier_epochs\"] == 500\n    assert PAPER_PROTOCOL[\"critic_weight_clip\"] == 0.01\n\n    clean_count = (\n        IMPLEMENTATION_ASSUMPTIONS[\"real_train_per_class\"]\n        + IMPLEMENTATION_ASSUMPTIONS[\"real_validation_per_class\"]\n        + IMPLEMENTATION_ASSUMPTIONS[\"real_test_per_class\"]\n    )\n    assert clean_count == PAPER_PROTOCOL[\"real_images_per_class\"]\n    assert IMPLEMENTATION_ASSUMPTIONS[\"real_test_per_class\"] == 5\n    print(\"Protocol and leakage-free split assertions passed.\")\n\nassert_protocol_configuration()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:23.477655Z","iopub.execute_input":"2026-07-14T15:51:23.477934Z","iopub.status.idle":"2026-07-14T15:51:23.483725Z","shell.execute_reply.started":"2026-07-14T15:51:23.477911Z","shell.execute_reply":"2026-07-14T15:51:23.483012Z"}},"outputs":[],"execution_count":null},{"id":"8de3b641","cell_type":"markdown","source":"## 2. Locate the two Kaggle datasets\n\nAttach:\n\n1. `chest-xray-covid19-pneumonia`\n2. `rsna-pneumonia-detection-challenge`\n\nThe first dataset supplies Normal, COVID-19, and Pneumonia images. The RSNA dataset supplies the Lung Opacity class using `Target = 1`, following the paper's dataset choice.","metadata":{}},{"id":"b93afb0e","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\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\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\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. Attach 'chest-xray-covid19-pneumonia'.\"\n)\nassert RSNA_ROOT is not None, (\n    \"RSNA dataset not found. Attach '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 = rsna_labels.groupby(\"patientId\", as_index=False)[\"Target\"].max()\nprint(rsna_unique[\"Target\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:23.484916Z","iopub.execute_input":"2026-07-14T15:51:23.485224Z","iopub.status.idle":"2026-07-14T15:51:39.404200Z","shell.execute_reply.started":"2026-07-14T15:51:23.485191Z","shell.execute_reply":"2026-07-14T15:51:39.403284Z"}},"outputs":[],"execution_count":null},{"id":"1e927b3a","cell_type":"markdown","source":"## 3. Select 25 real images per class and split them before training\n\nNo GAN, classifier, normalization statistic, or model-selection decision is allowed to use validation or test images.","metadata":{}},{"id":"967b3d97","cell_type":"code","source":"GAN_IMAGE_SIZE = PAPER_PROTOCOL[\"gan_image_size\"]\n\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\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\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\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\n\ndef chest_files(class_folder):\n    files = []\n    for original_split in (\"train\", \"test\"):\n        files.extend(glob.glob(str(CXR_ROOT / original_split / class_folder / \"*\")))\n    return sorted([f for f in files if Path(f).is_file()])\n\n\ndef choose_and_split(items, rng):\n    items = np.asarray(list(items), dtype=object)\n    required = PAPER_PROTOCOL[\"real_images_per_class\"]\n    if len(items) < required:\n        raise ValueError(f\"Need {required} items, but only {len(items)} are available.\")\n\n    chosen_indices = rng.choice(len(items), size=required, replace=False)\n    chosen = items[chosen_indices]\n    chosen = chosen[rng.permutation(len(chosen))]\n\n    n_train = IMPLEMENTATION_ASSUMPTIONS[\"real_train_per_class\"]\n    n_val = IMPLEMENTATION_ASSUMPTIONS[\"real_validation_per_class\"]\n    return {\n        \"train\": chosen[:n_train].tolist(),\n        \"validation\": chosen[n_train:n_train + n_val].tolist(),\n        \"test\": chosen[n_train + n_val:].tolist(),\n    }\n\n\nrng = np.random.default_rng(SEED)\nsource_pools = {\n    \"Normal\": chest_files(\"NORMAL\"),\n    \"COVID-19\": chest_files(\"COVID19\"),\n    \"Viral Pneumonia\": chest_files(\"PNEUMONIA\"),\n}\n\nopacity_ids = rsna_unique.loc[rsna_unique[\"Target\"] == 1, \"patientId\"].astype(str).tolist()\nsource_pools[\"Lung Opacity\"] = [\n    pid for pid in opacity_ids if (RSNA_IMAGE_DIR / f\"{pid}.dcm\").exists()\n]\n\nsplit_sources = {\n    class_name: choose_and_split(source_pools[class_name], rng)\n    for class_name in CLASSES\n}\n\nmanifest_rows = []\nfor class_name in CLASSES:\n    for split_name in (\"train\", \"validation\", \"test\"):\n        for source in split_sources[class_name][split_name]:\n            manifest_rows.append({\n                \"class\": class_name,\n                \"split\": split_name,\n                \"source\": str(source),\n            })\n\nsplit_manifest = pd.DataFrame(manifest_rows)\nsplit_manifest.to_csv(WORK_DIR / \"clean_real_split_manifest.csv\", index=False)\ndisplay(split_manifest.groupby([\"split\", \"class\"]).size().unstack(fill_value=0))\n\n# No source may occur in more than one split.\nassert split_manifest[\"source\"].nunique() == len(split_manifest)\nassert len(split_manifest) == 100\n\nsplit_signature = hashlib.sha256(\n    split_manifest.sort_values([\"class\", \"split\", \"source\"])\n    .to_csv(index=False)\n    .encode(\"utf-8\")\n).hexdigest()[:12]\nprint(\"Split signature:\", split_signature)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:39.405868Z","iopub.execute_input":"2026-07-14T15:51:39.406170Z","iopub.status.idle":"2026-07-14T15:51:41.914578Z","shell.execute_reply.started":"2026-07-14T15:51:39.406147Z","shell.execute_reply":"2026-07-14T15:51:41.913722Z"}},"outputs":[],"execution_count":null},{"id":"728a7d92","cell_type":"code","source":"def load_source(class_name, source):\n    if class_name == \"Lung Opacity\":\n        return load_rsna_image(source)\n    return load_chest_image(source)\n\n\nreal_images = {\"train\": {}, \"validation\": {}, \"test\": {}}\nfor split_name in real_images:\n    for class_name in CLASSES:\n        real_images[split_name][class_name] = np.stack([\n            load_source(class_name, source)\n            for source in split_sources[class_name][split_name]\n        ]).astype(np.uint8)\n        print(split_name, class_name, real_images[split_name][class_name].shape)\n\nassert all(len(real_images[\"train\"][c]) == 16 for c in CLASSES)\nassert all(len(real_images[\"validation\"][c]) == 4 for c in CLASSES)\nassert all(len(real_images[\"test\"][c]) == 5 for c in CLASSES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:41.915615Z","iopub.execute_input":"2026-07-14T15:51:41.915957Z","iopub.status.idle":"2026-07-14T15:51:45.609968Z","shell.execute_reply.started":"2026-07-14T15:51:41.915923Z","shell.execute_reply":"2026-07-14T15:51:45.609295Z"}},"outputs":[],"execution_count":null},{"id":"53a91a88","cell_type":"code","source":"fig, axes = plt.subplots(4, 3, figsize=(9, 11))\nfor row, class_name in enumerate(CLASSES):\n    for col, split_name in enumerate([\"train\", \"validation\", \"test\"]):\n        axes[row, col].imshow(real_images[split_name][class_name][0, ..., 0], cmap=\"gray\")\n        axes[row, col].set_title(f\"{class_name}\\n{split_name}\")\n        axes[row, col].axis(\"off\")\nplt.suptitle(\"Examples from the disjoint real-image splits\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:45.610873Z","iopub.execute_input":"2026-07-14T15:51:45.611175Z","iopub.status.idle":"2026-07-14T15:51:46.676471Z","shell.execute_reply.started":"2026-07-14T15:51:45.611151Z","shell.execute_reply":"2026-07-14T15:51:46.675669Z"}},"outputs":[],"execution_count":null},{"id":"a54d9b3e","cell_type":"markdown","source":"## 4. Assemble the fixed real train, validation, and test sets","metadata":{}},{"id":"7c39cffd","cell_type":"code","source":"def assemble_real_split(split_name):\n    X, y, origins = [], [], []\n    for class_id, class_name in enumerate(CLASSES):\n        images = real_images[split_name][class_name]\n        X.extend(images)\n        y.extend([class_id] * len(images))\n        origins.extend([\"real\"] * len(images))\n    return (\n        np.stack(X).astype(np.uint8),\n        np.asarray(y, dtype=np.int32),\n        np.asarray(origins),\n    )\n\n\nX_train_real, y_train_real, origin_train_real = assemble_real_split(\"train\")\nX_val_real, y_val_real, origin_val_real = assemble_real_split(\"validation\")\nX_test_real, y_test_real, origin_test_real = assemble_real_split(\"test\")\n\nprint(\"Real training  :\", X_train_real.shape, np.bincount(y_train_real))\nprint(\"Real validation:\", X_val_real.shape, np.bincount(y_val_real))\nprint(\"Real test      :\", X_test_real.shape, np.bincount(y_test_real))\n\nassert len(X_train_real) == 64\nassert len(X_val_real) == 16\nassert len(X_test_real) == 20\nassert np.all(origin_val_real == \"real\")\nassert np.all(origin_test_real == \"real\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:46.677713Z","iopub.execute_input":"2026-07-14T15:51:46.678033Z","iopub.status.idle":"2026-07-14T15:51:46.699494Z","shell.execute_reply.started":"2026-07-14T15:51:46.678008Z","shell.execute_reply":"2026-07-14T15:51:46.698788Z"}},"outputs":[],"execution_count":null},{"id":"913cc2dd","cell_type":"markdown","source":"## 5. Classifier preprocessing and transfer-learning models\n\nFeature-wise Z-normalization statistics are calculated only from the 64 real training images. The same statistics are used for the real-only and augmented experiments so that synthetic images cannot change the normalization of the validation or test evaluation.","metadata":{}},{"id":"0bbd031b","cell_type":"code","source":"AUTOTUNE = tf.data.AUTOTUNE\n\n\ndef compute_featurewise_z_stats(X_reference, 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_reference), batch_size):\n        batch = X_reference[start:start + batch_size].astype(np.float32)\n        resized = tf.image.resize(\n            batch,\n            (target_size, target_size),\n            method=\"bilinear\",\n            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    variance = np.maximum(total_sq / count - np.square(mean), 0.0)\n    std = np.sqrt(variance)\n    std[std < 1e-6] = 1.0\n    return mean.astype(np.float32), std.astype(np.float32)\n\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,\n            (target_size, target_size),\n            method=\"bilinear\",\n            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\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\",\n        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-14T15:51:46.701938Z","iopub.execute_input":"2026-07-14T15:51:46.702289Z","iopub.status.idle":"2026-07-14T15:51:46.712589Z","shell.execute_reply.started":"2026-07-14T15:51:46.702264Z","shell.execute_reply":"2026-07-14T15:51:46.711620Z"}},"outputs":[],"execution_count":null},{"id":"71bd16e7","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 attach \"\n            \"the matching Keras weight files. Random initialization would not match \"\n            \"the paper's transfer-learning setup.\"\n        ) from exc\n    raise ValueError(name)\n\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\n\ndef build_paper_classifier(name, class_count=NUM_CLASSES):\n    size = model_input_size(name)\n    backbone = load_imagenet_backbone(name, size)\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            IMPLEMENTATION_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    return keras.Model(inputs, outputs, name=f\"{name}_paper\"), backbone, size\n\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(y_true * tf.math.log(y_pred) * weights, axis=-1)\n\n    return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:46.713859Z","iopub.execute_input":"2026-07-14T15:51:46.714200Z","iopub.status.idle":"2026-07-14T15:51:46.729354Z","shell.execute_reply.started":"2026-07-14T15:51:46.714165Z","shell.execute_reply":"2026-07-14T15:51:46.728557Z"}},"outputs":[],"execution_count":null},{"id":"43695292","cell_type":"markdown","source":"## 6. Leakage-free OT-WGAN\n\nThe paper does not state whether the GAN is conditional or class-specific. This notebook uses one class-specific GAN per class, because it must generate class-labelled images for the four classifier categories.","metadata":{}},{"id":"3fecc28f","cell_type":"code","source":"GAN_SIZE = PAPER_PROTOCOL[\"gan_image_size\"]\nLATENT_DIM = IMPLEMENTATION_ASSUMPTIONS[\"gan_latent_dim\"]\nGAN_KERNEL = PAPER_PROTOCOL[\"gan_kernel_size\"]\nGAN_STRIDE = PAPER_PROTOCOL[\"gan_stride\"]\n\n\ndef build_paper_generator():\n    start = PAPER_PROTOCOL[\"generator_start_size\"]\n    c0 = IMPLEMENTATION_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\"generator_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    return keras.Model(z, output, name=\"paper_generator\")\n\n\ndef build_paper_critic():\n    image = layers.Input((GAN_SIZE, GAN_SIZE, 1), name=\"image_512\")\n    x = image\n\n    # The paper says Conv2DTranspose while also saying 512 -> 64 downsampling.\n    # A stride-2 transpose convolution upsamples, so Conv2D is used to preserve\n    # the stated downsampling geometry.\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\n_test_generator = build_paper_generator()\n_test_critic = build_paper_critic()\nprint(\"Generator:\", _test_generator.input_shape, \"->\", _test_generator.output_shape)\nprint(\"Critic   :\", _test_critic.input_shape, \"->\", _test_critic.output_shape)\nassert _test_generator.output_shape[1:3] == (512, 512)\nassert tuple(_test_critic.get_layer(\"downsample_to_64\").output.shape[1:3]) == (64, 64)\ndel _test_generator, _test_critic\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:46.730347Z","iopub.execute_input":"2026-07-14T15:51:46.730751Z","iopub.status.idle":"2026-07-14T15:51:48.408883Z","shell.execute_reply.started":"2026-07-14T15:51:46.730716Z","shell.execute_reply":"2026-07-14T15:51:48.408189Z"}},"outputs":[],"execution_count":null},{"id":"3c2649e6","cell_type":"markdown","source":"### Debiased Sinkhorn divergence\n\nThe earlier notebook calculated only a one-way entropy-regularized transport cost. The function below subtracts the two self-transport terms, which is the defining correction in Sinkhorn divergence.","metadata":{}},{"id":"8f57573c","cell_type":"code","source":"def pooled_features(images):\n    pool_size = IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_pool_size\"]\n    assert GAN_SIZE % pool_size == 0\n    stride = GAN_SIZE // pool_size\n    pooled = tf.nn.avg_pool2d(\n        images,\n        ksize=stride,\n        strides=stride,\n        padding=\"VALID\",\n    )\n    return tf.reshape(pooled, (tf.shape(pooled)[0], -1))\n\n\ndef pairwise_squared_cost(x, y):\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 = x2 + tf.transpose(y2) - 2.0 * tf.matmul(x, y, transpose_b=True)\n    return tf.maximum(cost, 0.0) / feature_count\n\n\ndef entropy_regularized_ot_value(x, y):\n    cost = pairwise_squared_cost(x, y)\n    epsilon = tf.constant(\n        IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_epsilon\"],\n        dtype=tf.float32,\n    )\n\n    n = tf.shape(x)[0]\n    m = tf.shape(y)[0]\n    log_a = -tf.math.log(tf.cast(n, tf.float32))\n    log_b = -tf.math.log(tf.cast(m, tf.float32))\n\n    f = tf.zeros((n, 1), dtype=tf.float32)\n    g = tf.zeros((m, 1), dtype=tf.float32)\n\n    for _ in range(IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_iterations\"]):\n        f = epsilon * (\n            log_a\n            - tf.reduce_logsumexp(\n                (tf.transpose(g) - cost) / epsilon,\n                axis=1,\n                keepdims=True,\n            )\n        )\n        g = epsilon * tf.transpose(\n            log_b\n            - tf.reduce_logsumexp(\n                (f - cost) / epsilon,\n                axis=0,\n                keepdims=True,\n            )\n        )\n\n    log_transport = (f + tf.transpose(g) - cost) / epsilon\n    safe_log_transport = tf.clip_by_value(log_transport, -60.0, 20.0)\n    transport = tf.exp(safe_log_transport)\n\n    transport_cost = tf.reduce_sum(transport * cost)\n    negative_entropy = tf.reduce_sum(\n        transport * (safe_log_transport - 1.0)\n    )\n    return transport_cost + epsilon * negative_entropy\n\n\ndef sinkhorn_divergence(real, fake):\n    x = pooled_features(real)\n    y = pooled_features(fake)\n    ot_xy = entropy_regularized_ot_value(x, y)\n    ot_xx = entropy_regularized_ot_value(x, x)\n    ot_yy = entropy_regularized_ot_value(y, y)\n    return ot_xy - 0.5 * ot_xx - 0.5 * ot_yy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:48.409752Z","iopub.execute_input":"2026-07-14T15:51:48.410143Z","iopub.status.idle":"2026-07-14T15:51:48.420542Z","shell.execute_reply.started":"2026-07-14T15:51:48.410120Z","shell.execute_reply":"2026-07-14T15:51:48.419802Z"}},"outputs":[],"execution_count":null},{"id":"8b5ad05d","cell_type":"code","source":"class LeakageFreeOTWGANTrainer:\n    def __init__(self):\n        self.generator = build_paper_generator()\n        self.critic = build_paper_critic()\n        self.generator_optimizer = keras.optimizers.Adam()\n        self.critic_optimizer = keras.optimizers.Adam()\n        self.gradient_clip_norm = IMPLEMENTATION_ASSUMPTIONS[\"gradient_clip_norm\"]\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 = tf.stop_gradient(self.generator(z, training=True))\n            real_score = self.critic(real, training=True)\n            fake_score = self.critic(fake, training=True)\n            critic_loss = tf.reduce_mean(fake_score) - tf.reduce_mean(real_score)\n\n        gradients = tape.gradient(critic_loss, self.critic.trainable_variables)\n        gradients, _ = tf.clip_by_global_norm(gradients, self.gradient_clip_norm)\n        self.critic_optimizer.apply_gradients(\n            zip(gradients, self.critic.trainable_variables)\n        )\n\n        clip_value = PAPER_PROTOCOL[\"critic_weight_clip\"]\n        for variable in self.critic.trainable_variables:\n            variable.assign(tf.clip_by_value(variable, -clip_value, clip_value))\n\n        wasserstein_estimate = tf.reduce_mean(real_score) - tf.reduce_mean(fake_score)\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            adversarial_loss = -tf.reduce_mean(fake_score)\n            sinkhorn_loss = sinkhorn_divergence(real, fake)\n            generator_loss = (\n                adversarial_loss\n                + IMPLEMENTATION_ASSUMPTIONS[\"sinkhorn_weight\"] * sinkhorn_loss\n            )\n\n        gradients = tape.gradient(generator_loss, self.generator.trainable_variables)\n        gradients, _ = tf.clip_by_global_norm(gradients, self.gradient_clip_norm)\n        self.generator_optimizer.apply_gradients(\n            zip(gradients, self.generator.trainable_variables)\n        )\n        return generator_loss, adversarial_loss, sinkhorn_loss\n\n    def fit(self, class_name, real_train_uint8):\n        # Only this class's real training images are passed here.\n        real = real_train_uint8.astype(np.float32)\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 = IMPLEMENTATION_ASSUMPTIONS[\"gan_batch_size\"]\n        critic_updates = IMPLEMENTATION_ASSUMPTIONS[\"critic_updates_per_generator\"]\n        log_every = IMPLEMENTATION_ASSUMPTIONS[\"gan_log_every\"]\n\n        history_rows = []\n        for epoch in range(GAN_EPOCHS):\n            for _ in range(critic_updates):\n                indices = np.random.choice(len(real_z), size=batch_size, replace=True)\n                critic_loss, w_est = self.critic_step(\n                    tf.convert_to_tensor(real_z[indices], dtype=tf.float32)\n                )\n\n            indices = np.random.choice(len(real_z), size=batch_size, replace=True)\n            g_loss, adversarial_loss, sinkhorn_loss = self.generator_step(\n                tf.convert_to_tensor(real_z[indices], dtype=tf.float32)\n            )\n\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_adversarial_loss\": float(adversarial_loss),\n                \"sinkhorn_divergence\": float(sinkhorn_loss),\n            }\n            history_rows.append(row)\n\n            if (epoch + 1) % log_every == 0 or epoch == 0 or epoch == GAN_EPOCHS - 1:\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\"S={row['sinkhorn_divergence']:.5f}\"\n                )\n\n        return pd.DataFrame(history_rows), mean.astype(np.float32), std.astype(np.float32)\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,"execution":{"iopub.status.busy":"2026-07-14T15:51:48.421614Z","iopub.execute_input":"2026-07-14T15:51:48.421927Z","iopub.status.idle":"2026-07-14T15:51:48.441656Z","shell.execute_reply.started":"2026-07-14T15:51:48.421895Z","shell.execute_reply":"2026-07-14T15:51:48.440850Z"}},"outputs":[],"execution_count":null},{"id":"9e50165a","cell_type":"markdown","source":"## 7. Train the GANs only on real training images and generate 1,600 training images\n\nNo validation or test image is passed to any GAN.","metadata":{}},{"id":"79cd7bde","cell_type":"code","source":"synthetic_images = {}\ngan_history_files = {}\ngan_quality_rows = []\n\nfor class_name in CLASSES:\n    safe_name = class_name.lower().replace(\" \", \"_\").replace(\"-\", \"\")\n    cache_path = WORK_DIR / \"gan_cache\" / f\"{safe_name}_{split_signature}_synthetic.npz\"\n    history_path = WORK_DIR / \"gan_cache\" / f\"{safe_name}_{split_signature}_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 clean-split cache:\", synthetic_images[class_name].shape)\n        gan_history_files[class_name] = str(history_path)\n    else:\n        if not RUN_CONTROL[\"run_gan_training\"]:\n            raise FileNotFoundError(\n                f\"GAN training is disabled, but clean-split cache is missing: {cache_path}\"\n            )\n\n        print(\"\\n\" + \"=\" * 80)\n        print(\"Training OT-WGAN using TRAIN images only:\", class_name)\n        print(\"Number of real images seen by GAN:\", len(real_images[\"train\"][class_name]))\n\n        tf.keras.backend.clear_session()\n        gc.collect()\n\n        trainer = LeakageFreeOTWGANTrainer()\n        history, gan_mean, gan_std = trainer.fit(\n            class_name,\n            real_images[\"train\"][class_name],\n        )\n        generated = trainer.generate(\n            PAPER_PROTOCOL[\"generated_images_per_class\"],\n            gan_mean,\n            gan_std,\n        )\n\n        np.savez_compressed(\n            cache_path,\n            images=generated,\n            split_signature=split_signature,\n            class_name=class_name,\n        )\n        history.to_csv(history_path, index=False)\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()\n\n    images = synthetic_images[class_name]\n    flat_std = images.reshape(len(images), -1).astype(np.float32).std(axis=1)\n    gan_quality_rows.append({\n        \"Class\": class_name,\n        \"Generated\": len(images),\n        \"MeanPixelStd\": float(flat_std.mean()),\n        \"FractionPixelStdAbove3\": float(np.mean(flat_std > 3.0)),\n    })\n\nassert sum(len(synthetic_images[c]) for c in CLASSES) == 1600\nassert all(len(synthetic_images[c]) == 400 for c in CLASSES)\n\ngan_quality_df = pd.DataFrame(gan_quality_rows)\ndisplay(gan_quality_df.round(4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-14T15:51:48.442840Z","iopub.execute_input":"2026-07-14T15:51:48.443198Z"}},"outputs":[],"execution_count":null},{"id":"495fd6cc","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(real_images[\"train\"][class_name][0, ..., 0], cmap=\"gray\")\n    axes[row, 0].set_title(\"Real train\")\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(synthetic_images[class_name][col - 1, ..., 0], cmap=\"gray\")\n        axes[row, col].set_title(\"Generated\")\n        axes[row, col].axis(\"off\")\n\nplt.suptitle(\"Real training examples versus generated training images\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"b712d251","cell_type":"markdown","source":"## 8. Build the augmented training set\n\nOnly the training set is augmented. The validation and test arrays remain unchanged and real-only.","metadata":{}},{"id":"acfa816b","cell_type":"code","source":"def assemble_augmented_training_set():\n    X_parts = [X_train_real]\n    y_parts = [y_train_real]\n    origin_parts = [origin_train_real]\n\n    for class_id, class_name in enumerate(CLASSES):\n        generated = synthetic_images[class_name]\n        X_parts.append(generated)\n        y_parts.append(np.full(len(generated), class_id, dtype=np.int32))\n        origin_parts.append(np.full(len(generated), \"synthetic\", dtype=\"<U9\"))\n\n    X = np.concatenate(X_parts, axis=0).astype(np.uint8)\n    y = np.concatenate(y_parts, axis=0).astype(np.int32)\n    origins = np.concatenate(origin_parts, axis=0)\n    return X, y, origins\n\n\nX_train_augmented, y_train_augmented, origin_train_augmented = assemble_augmented_training_set()\n\nprint(\"Augmented training shape:\", X_train_augmented.shape)\nprint(\"Augmented class counts:\", np.bincount(y_train_augmented))\nprint(\"Augmented origin counts:\", dict(zip(\n    *np.unique(origin_train_augmented, return_counts=True)\n)))\n\nassert len(X_train_augmented) == 1664\nassert np.all(np.bincount(y_train_augmented) == 416)\nassert np.all(origin_val_real == \"real\")\nassert np.all(origin_test_real == \"real\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"df5a789b","cell_type":"markdown","source":"## 9. Train and evaluate the real-only and augmented classifiers\n\nBoth stages use the same real validation set and the same unseen real test set. A best-validation-loss checkpoint is saved, but training still runs for the paper's stated 500 epochs unless quick-debug mode is enabled.","metadata":{}},{"id":"65252603","cell_type":"code","source":"def evaluate_predictions(stage, model_name, y_true, y_prob):\n    y_pred = np.argmax(y_prob, axis=1)\n    accuracy = accuracy_score(y_true, y_pred)\n    precision, recall, f1, _ = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=np.arange(NUM_CLASSES),\n        zero_division=0,\n    )\n\n    rows = []\n    for class_id, class_name in enumerate(CLASSES):\n        binary_truth = (y_true == class_id).astype(np.int32)\n        fpr, tpr, _ = roc_curve(binary_truth, y_prob[:, class_id])\n        rows.append({\n            \"Stage\": stage,\n            \"Model\": model_name,\n            \"Class\": class_name,\n            \"Accuracy\": accuracy,\n            \"Precision_PPV\": precision[class_id],\n            \"Recall_Sensitivity\": recall[class_id],\n            \"F1\": f1[class_id],\n            \"AUC\": auc(fpr, tpr),\n        })\n    return rows, y_pred\n\n\nall_metric_rows = []\nexperiment_index = []\n\n\ndef run_classifier_experiment(stage, X_train, y_train, train_origins):\n    assert np.all(origin_val_real == \"real\")\n    assert np.all(origin_test_real == \"real\")\n\n    print(\"\\n\" + \"=\" * 80)\n    print(\"Stage:\", stage)\n    print(\"Train origins:\", dict(zip(*np.unique(train_origins, return_counts=True))))\n    print(\"Validation origins: real only; n =\", len(X_val_real))\n    print(\"Test origins: real only; n =\", len(X_test_real))\n\n    class_weights = compute_class_weight(\n        class_weight=\"balanced\",\n        classes=np.arange(NUM_CLASSES),\n        y=y_train,\n    )\n    print(\"Class weights:\", dict(zip(CLASSES, class_weights)))\n\n    stage_records = []\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(model_name)\n\n        # Fixed reference: real training images only, shared across both stages.\n        mean, std = compute_featurewise_z_stats(X_train_real, input_size)\n\n        train_ds = make_classifier_dataset(\n            X_train,\n            y_train,\n            input_size,\n            NUM_CLASSES,\n            mean,\n            std,\n            training=True,\n        )\n        val_ds = make_classifier_dataset(\n            X_val_real,\n            y_val_real,\n            input_size,\n            NUM_CLASSES,\n            mean,\n            std,\n            training=False,\n        )\n        test_ds = make_classifier_dataset(\n            X_test_real,\n            y_test_real,\n            input_size,\n            NUM_CLASSES,\n            mean,\n            std,\n            training=False,\n        )\n\n        model.compile(\n            optimizer=keras.optimizers.Adam(),\n            loss=weighted_categorical_crossentropy(class_weights),\n            metrics=[\"accuracy\"],\n        )\n\n        tag = f\"{stage}__multiclass__{model_name}\"\n        history_csv = WORK_DIR / \"histories\" / f\"{tag}.csv\"\n        model_path = WORK_DIR / \"models\" / f\"{tag}.weights.h5\"\n\n        callbacks = [\n            keras.callbacks.CSVLogger(history_csv),\n            keras.callbacks.TerminateOnNaN(),\n        ]\n        if IMPLEMENTATION_ASSUMPTIONS[\"use_best_validation_checkpoint\"]:\n            callbacks.append(\n                keras.callbacks.ModelCheckpoint(\n                    model_path,\n                    monitor=\"val_loss\",\n                    save_best_only=True,\n                    save_weights_only=True,\n                    verbose=1,\n                )\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        val_losses = np.asarray(history.history[\"val_loss\"], dtype=np.float64)\n        best_epoch = int(np.nanargmin(val_losses) + 1)\n\n        if IMPLEMENTATION_ASSUMPTIONS[\"use_best_validation_checkpoint\"]:\n            model.load_weights(model_path)\n        else:\n            model.save_weights(model_path)\n\n        y_prob = model.predict(test_ds, verbose=1)\n        metric_rows, y_pred = evaluate_predictions(\n            stage,\n            model_name,\n            y_test_real,\n            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=y_test_real,\n            y_prob=y_prob,\n            y_pred=y_pred,\n            X_test=X_test_real,\n            origin_test=origin_test_real,\n            class_names=np.asarray(CLASSES),\n        )\n\n        stats_path = WORK_DIR / \"predictions\" / f\"{tag}_zstats.npz\"\n        np.savez_compressed(stats_path, mean=mean, std=std)\n\n        stage_records.append({\n            \"Stage\": stage,\n            \"Model\": model_name,\n            \"InputSize\": input_size,\n            \"Epochs\": CLASSIFIER_EPOCHS,\n            \"BestValidationEpoch\": best_epoch,\n            \"TrainReal\": int(np.sum(train_origins == \"real\")),\n            \"TrainSynthetic\": int(np.sum(train_origins == \"synthetic\")),\n            \"ValidationReal\": len(X_val_real),\n            \"TestReal\": len(X_test_real),\n            \"HistoryCSV\": str(history_csv),\n            \"PredictionFile\": str(pred_path),\n            \"ZStatsFile\": str(stats_path),\n            \"ModelFile\": str(model_path),\n        })\n\n        print(classification_report(\n            y_test_real,\n            y_pred,\n            target_names=CLASSES,\n            zero_division=0,\n        ))\n\n        del model, backbone, train_ds, val_ds, test_ds, y_prob, history\n        gc.collect()\n\n    return stage_records","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"c7387e8d","cell_type":"code","source":"if RUN_CONTROL[\"run_real_only_baseline\"]:\n    experiment_index.extend(\n        run_classifier_experiment(\n            \"real_only\",\n            X_train_real,\n            y_train_real,\n            origin_train_real,\n        )\n    )\nelse:\n    print(\"Real-only baseline skipped by RUN_CONTROL.\")\n\nif RUN_CONTROL[\"run_augmented_models\"]:\n    experiment_index.extend(\n        run_classifier_experiment(\n            \"real_plus_synthetic\",\n            X_train_augmented,\n            y_train_augmented,\n            origin_train_augmented,\n        )\n    )\nelse:\n    print(\"Augmented classifier stage skipped by RUN_CONTROL.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"82ac6a35","cell_type":"markdown","source":"## 10. Leakage audit\n\nThis cell must pass before the results are presented.","metadata":{}},{"id":"17bd7bb1","cell_type":"code","source":"train_source_set = set(split_manifest.loc[split_manifest[\"split\"] == \"train\", \"source\"])\nvalidation_source_set = set(split_manifest.loc[split_manifest[\"split\"] == \"validation\", \"source\"])\ntest_source_set = set(split_manifest.loc[split_manifest[\"split\"] == \"test\", \"source\"])\n\nassert train_source_set.isdisjoint(validation_source_set)\nassert train_source_set.isdisjoint(test_source_set)\nassert validation_source_set.isdisjoint(test_source_set)\nassert np.all(origin_val_real == \"real\")\nassert np.all(origin_test_real == \"real\")\nassert not np.any(origin_train_real == \"synthetic\")\nassert np.sum(origin_train_augmented == \"synthetic\") == 1600\n\nleakage_audit = pd.DataFrame([\n    {\"Check\": \"Real source overlap: train vs validation\", \"Passed\": train_source_set.isdisjoint(validation_source_set)},\n    {\"Check\": \"Real source overlap: train vs test\", \"Passed\": train_source_set.isdisjoint(test_source_set)},\n    {\"Check\": \"Real source overlap: validation vs test\", \"Passed\": validation_source_set.isdisjoint(test_source_set)},\n    {\"Check\": \"Validation contains only real images\", \"Passed\": bool(np.all(origin_val_real == \"real\"))},\n    {\"Check\": \"Test contains only real images\", \"Passed\": bool(np.all(origin_test_real == \"real\"))},\n    {\"Check\": \"All 1,600 synthetic images are training-only\", \"Passed\": int(np.sum(origin_train_augmented == \"synthetic\")) == 1600},\n])\ndisplay(leakage_audit)\nassert leakage_audit[\"Passed\"].all()\nprint(\"Leakage audit passed.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"fad518ad","cell_type":"markdown","source":"## 11. Result tables","metadata":{}},{"id":"c9947720","cell_type":"code","source":"metrics_df = pd.DataFrame(all_metric_rows)\nindex_df = pd.DataFrame(experiment_index)\n\nmetrics_path = WORK_DIR / \"leakage_free_metrics.csv\"\nindex_path = WORK_DIR / \"experiment_index.csv\"\nmetrics_df.to_csv(metrics_path, index=False)\nindex_df.to_csv(index_path, index=False)\n\nif len(metrics_df):\n    display(metrics_df.round(4))\n    model_summary = (\n        metrics_df.groupby([\"Stage\", \"Model\"], sort=False)\n        .agg(\n            Accuracy=(\"Accuracy\", \"first\"),\n            MacroPrecision=(\"Precision_PPV\", \"mean\"),\n            MacroRecall=(\"Recall_Sensitivity\", \"mean\"),\n            MacroF1=(\"F1\", \"mean\"),\n            MacroAUC=(\"AUC\", \"mean\"),\n        )\n        .reset_index()\n    )\n    display(model_summary.round(4))\n\n    accuracy_comparison = model_summary.pivot(\n        index=\"Model\",\n        columns=\"Stage\",\n        values=\"Accuracy\",\n    )\n    if {\"real_only\", \"real_plus_synthetic\"}.issubset(accuracy_comparison.columns):\n        accuracy_comparison[\"AugmentationChange\"] = (\n            accuracy_comparison[\"real_plus_synthetic\"]\n            - accuracy_comparison[\"real_only\"]\n        )\n    display(accuracy_comparison.round(4))\nelse:\n    print(\"No classifier results were produced.\")\n\ndisplay(index_df)\nprint(\"Saved:\", metrics_path)\nprint(\"Saved:\", index_path)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"a8ab76d8","cell_type":"markdown","source":"## 12. Training curves","metadata":{}},{"id":"603b7691","cell_type":"code","source":"for _, row in index_df.iterrows():\n    history = pd.read_csv(row[\"HistoryCSV\"])\n    plt.figure(figsize=(7, 4))\n    plt.plot(history[\"epoch\"], history[\"accuracy\"], label=\"Train accuracy\")\n    plt.plot(history[\"epoch\"], history[\"val_accuracy\"], label=\"Real validation accuracy\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"Accuracy\")\n    plt.title(f\"{row['Stage']} | {row['Model']}\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"0fff44d3","cell_type":"markdown","source":"## 13. Confusion matrices and ROC curves on unseen real test images","metadata":{}},{"id":"e77d9f6a","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    origin_test = pred[\"origin_test\"]\n    assert np.all(origin_test == \"real\")\n\n    cm = confusion_matrix(y_true, y_pred, normalize=\"true\")\n    plt.figure(figsize=(6, 5))\n    plt.imshow(cm, interpolation=\"nearest\")\n    plt.title(f\"Real test confusion matrix | {row['Stage']} | {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    for i in range(cm.shape[0]):\n        for j in range(cm.shape[1]):\n            plt.text(j, i, f\"{cm[i, j]:.2f}\", ha=\"center\", va=\"center\")\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\"Real test ROC | {row['Stage']} | {row['Model']}\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"18321312","cell_type":"markdown","source":"## 14. Grad-CAM on unseen real test images","metadata":{}},{"id":"2a5dcbbf","cell_type":"code","source":"def gradcam_heatmap(model, preprocessed_image):\n    grad_model = keras.Model(\n        model.inputs,\n        [model.get_layer(\"last_conv_features\").output, model.output],\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        predicted_class = tf.argmax(predictions[0])\n        class_score = predictions[:, predicted_class]\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(predicted_class.numpy())\n\n\nfor _, row in index_df.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    origin_test = pred[\"origin_test\"]\n    assert np.all(origin_test == \"real\")\n\n    model, _, _ = build_paper_classifier(row[\"Model\"])\n    model.load_weights(row[\"ModelFile\"])\n\n    input_size = int(row[\"InputSize\"])\n    mean = stats[\"mean\"]\n    std = stats[\"std\"]\n\n    fig, axes = plt.subplots(1, NUM_CLASSES, figsize=(14, 4))\n    for class_id, class_name in enumerate(CLASSES):\n        index = np.where(y_true == class_id)[0][0]\n        prepared = z_normalize_one(X_test[index], input_size, mean, std)\n        heatmap, predicted_class = gradcam_heatmap(model, prepared)\n\n        heatmap_resized = tf.image.resize(\n            heatmap[..., None],\n            (input_size, input_size),\n        ).numpy().squeeze()\n        original_resized = tf.image.resize(\n            X_test[index].astype(np.float32),\n            (input_size, input_size),\n            method=\"bilinear\",\n            antialias=True,\n        ).numpy().squeeze()\n\n        axes[class_id].imshow(original_resized, cmap=\"gray\")\n        axes[class_id].imshow(heatmap_resized, alpha=0.42)\n        axes[class_id].set_title(\n            f\"True: {class_name}\\nPred: {CLASSES[predicted_class]}\"\n        )\n        axes[class_id].axis(\"off\")\n\n    fig.suptitle(f\"Real-test Grad-CAM | {row['Stage']} | {row['Model']}\")\n    fig.tight_layout()\n    plt.show()\n\n    del model\n    gc.collect()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"eddd7e35","cell_type":"markdown","source":"## 15. How to report the result\n\nReport the real-only and augmented results together. The augmentation method is supported only when the augmented model improves performance on the same unseen real test images. Do not use synthetic test accuracy as evidence of medical classification performance.\n\nThe paper reports VGG19 96.94%, InceptionV3 95.35%, and Xception 92.62%, but this corrected evaluation is not expected to reproduce those exact values because the paper does not disclose its exact split, GAN settings, or whether synthetic images were included in testing.","metadata":{}}]}