{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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":[{"cell_type":"markdown","source":"# SSL-Conditioned GAN for Chest X-Ray Synthesis\nGenerates synthetic CXR images conditioned on **pre-computed SSL embeddings** from DINOv2:\n`image_embeddings_128d` ⊕ `disease_onehot` ⊕ `severity_onehot` ⊕ `demographic_embeddings_64d`.\n\nAll conditioning vectors are loaded from `.npy` files produced by the SSL pipeline — no raw dataset access required during GAN training.","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport gc\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport torchvision.utils as vutils\n\n# -------------------------------------------------------------------\n# Architecture & training constants\n# -------------------------------------------------------------------\nIMAGE_SIZE  = 256\nCHANNELS    = 1\nBATCH_SIZE  = 32       # reduced from 64 to save VRAM\nNUM_WORKERS = 4\n\nZ_DIM = 128\n\n# SSL conditioning sub-dimensions (match SSL notebook outputs)\nSSL_EMB_DIM         = 128\nDISEASE_ONEHOT_DIM  = 3     # updated after loading\nSEVERITY_ONEHOT_DIM = 4\nDEMO_EMB_DIM        = 64\n\n# Total conditioning dim -- recomputed after npy load\nCOND_DIM = SSL_EMB_DIM + DISEASE_ONEHOT_DIM + SEVERITY_ONEHOT_DIM + DEMO_EMB_DIM  # 199\n\n# Loss weights\nLAMBDA_CLS  = 1.0\nLAMBDA_SEV  = 0.5\nLAMBDA_DEMO = 0.5\nLAMBDA_SSL  = 2.0    # reduced from 5.0\nSSL_EVERY   = 4      # compute DINOv2 perceptual loss every N batches only\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Using device: {device}')\nif torch.cuda.is_available():\n    print(f'GPU(s): {torch.cuda.device_count()} x {torch.cuda.get_device_name(0)}')\nprint(f'Initial COND_DIM: {COND_DIM}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:51:42.614219Z","iopub.execute_input":"2026-06-22T06:51:42.615074Z","iopub.status.idle":"2026-06-22T06:51:42.622985Z","shell.execute_reply.started":"2026-06-22T06:51:42.615023Z","shell.execute_reply":"2026-06-22T06:51:42.622066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Paths -- update NPY_DIR to wherever final_conditions/ was saved\n# -------------------------------------------------------------------\nNPY_DIR        = Path('/kaggle/input/datasets/kaushik2k25/embeddings')\nPARQUET_PATH   = '/kaggle/input/datasets/kaushik2k25/finalparquet/unchanged_train.parquet'\n\nCACHE_DIR       = Path('/kaggle/working/dcm_cache')\nOUTPUT_GRID_DIR = Path('/kaggle/working/ssl_gan_grids')\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\nOUTPUT_GRID_DIR.mkdir(parents=True, exist_ok=True)\n\nDATASET_ROOTS = {\n    'CheXpert'  : '/kaggle/input/datasets/willarevalo/chexpert-v10-small/CheXpert-v1.0-small',\n    'Pediatric' : '/kaggle/input/datasets/paultimothymooney/chest-xray-pneumonia/chest_xray',\n    'NIH'       : '/kaggle/input/datasets/organizations/nih-chest-xrays/data',\n    'COVIDx'    : '/kaggle/input/datasets/andyczhao/covidx-cxr2',\n    'RSNA'      : '/kaggle/input/competitions/rsna-pneumonia-detection-challenge/stage_2_train_images'\n}\n\nCLASS_NAMES = ['No Finding', 'Pneumonia', 'COVID-19']\nDISEASE_MAP = {'No Finding': 0, 'Normal': 0, 'Pneumonia': 1, 'COVID-19': 2}\n\nprint('Paths configured.')\nprint(f'  NPY dir  : {NPY_DIR}')\nprint(f'  Parquet  : {PARQUET_PATH}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:51:44.909848Z","iopub.execute_input":"2026-06-22T06:51:44.910612Z","iopub.status.idle":"2026-06-22T06:51:44.916670Z","shell.execute_reply.started":"2026-06-22T06:51:44.910556Z","shell.execute_reply":"2026-06-22T06:51:44.915805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 3 -- Load Pre-computed SSL Conditioning Arrays\n# -------------------------------------------------------------------\nprint('Loading SSL conditioning arrays...')\n\nssl_embeddings  = np.load(NPY_DIR / 'features_128d_full.npy').astype(np.float32)    # (N, 128)\ndisease_onehot  = np.load(NPY_DIR / 'disease_onehot_full.npy').astype(np.float32)           # (N, D)\nseverity_onehot = np.load(NPY_DIR / 'severity_onehot_full.npy').astype(np.float32)          # (N, 4)\ndemo_embeddings = np.load(NPY_DIR / 'demographic_embeddings_64d_full.npy').astype(np.float32)  # (N, 64)\n\n# Recompute actual dims after loading (disease column count can vary)\nDISEASE_ONEHOT_DIM  = disease_onehot.shape[1]\nSEVERITY_ONEHOT_DIM = severity_onehot.shape[1]\nCOND_DIM = SSL_EMB_DIM + DISEASE_ONEHOT_DIM + SEVERITY_ONEHOT_DIM + DEMO_EMB_DIM\n\n# Concatenate into a single fused conditioning matrix: (N, COND_DIM)\ncond_matrix = np.concatenate(\n    [ssl_embeddings, disease_onehot, severity_onehot, demo_embeddings], axis=1\n).astype(np.float32)\n\nN_TOTAL = cond_matrix.shape[0]\n\nprint(f'  ssl_embeddings  : {ssl_embeddings.shape}')\nprint(f'  disease_onehot  : {disease_onehot.shape}')\nprint(f'  severity_onehot : {severity_onehot.shape}')\nprint(f'  demo_embeddings : {demo_embeddings.shape}')\nprint(f'  cond_matrix     : {cond_matrix.shape}  (fused)')\nprint(f'  COND_DIM        : {COND_DIM}')\nprint(f'  N_TOTAL         : {N_TOTAL:,} samples')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:52:47.945341Z","iopub.execute_input":"2026-06-22T06:52:47.946031Z","iopub.status.idle":"2026-06-22T06:52:48.789305Z","shell.execute_reply.started":"2026-06-22T06:52:47.945990Z","shell.execute_reply":"2026-06-22T06:52:48.788584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 4 -- Load Parquet & Resolve Absolute Paths\n# The parquet provides image paths that map 1-to-1 with npy rows\n# (same ordering produced by the SSL pipeline).\n# -------------------------------------------------------------------\ndef resolve_path(row):\n    \"\"\"Maps a parquet-relative path to its absolute Kaggle dataset path.\"\"\"\n    ds_name = row.get('dataset')\n    root    = DATASET_ROOTS.get(ds_name)\n    if root is None:\n        return str(row['path'])\n    p = str(row['path']).lstrip('/')\n    if ds_name == 'RSNA' and not p.endswith('.dcm'):\n        return str(Path(root) / (Path(p).stem + '.dcm'))\n    return str(Path(root) / p)\n\n\nprint(f'Loading parquet: {PARQUET_PATH}')\ndf_raw = pd.read_parquet(PARQUET_PATH)\ndf_raw = df_raw.dropna(subset=['disease']).reset_index(drop=True)\ndf_raw['abs_path'] = df_raw.apply(resolve_path, axis=1)\n\nprint(f'  Parquet rows : {len(df_raw):,}')\nprint(f'  NPY rows     : {N_TOTAL:,}')\n\n# Truncate to N_TOTAL first to stay aligned with npy arrays\nif len(df_raw) > N_TOTAL:\n    df_raw = df_raw.iloc[:N_TOTAL].reset_index(drop=True)\nelif len(df_raw) < N_TOTAL:\n    raise ValueError(\n        f'Parquet has fewer rows ({len(df_raw)}) than npy arrays ({N_TOTAL}). '\n        f'Re-run the SSL pipeline against this parquet.'\n    )\n\n# ── Cap per class using original indices so npy alignment is preserved\nMAX_PER_CLASS = 500\nsampled_indices = []\nfor disease_label, group in df_raw.groupby('disease'):\n    sampled_indices.extend(\n        group.sample(n=min(len(group), MAX_PER_CLASS), random_state=42).index.tolist()\n    )\nsampled_indices = sorted(sampled_indices)   # keep ascending order for npy slicing\n\ndf_raw = df_raw.loc[sampled_indices].reset_index(drop=True)\ncond_matrix_sampled = cond_matrix[sampled_indices]   # slice npy rows by same indices\n\nprint(f'  After capping ({MAX_PER_CLASS}/class): {len(df_raw):,} rows')\n\n# Verify file existence\nprint('Verifying file paths...')\nvalid_mask     = [Path(row['abs_path']).exists() for _, row in df_raw.iterrows()]\nvalid_positions = [i for i, v in enumerate(valid_mask) if v]\n\ndf_train   = df_raw.iloc[valid_positions].reset_index(drop=True)\ncond_train = cond_matrix_sampled[valid_positions]   # aligned slice\n\nprint(f'  Valid samples : {len(df_train):,}')\nprint(f'  cond_train    : {cond_train.shape}')\nprint(f'  Batches/epoch : {len(df_train) // BATCH_SIZE}')\nprint('\\n--- Disease Distribution ---')\nprint(df_train['disease'].value_counts())\nprint('----------------------------')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:52:53.308797Z","iopub.execute_input":"2026-06-22T06:52:53.309407Z","iopub.status.idle":"2026-06-22T06:57:38.330227Z","shell.execute_reply.started":"2026-06-22T06:52:53.309374Z","shell.execute_reply":"2026-06-22T06:57:38.329537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 5 -- Image Loading Utilities (DICOM + CLAHE)\n# Identical to original GAN notebook.\n# -------------------------------------------------------------------\ndef apply_clahe(img_array: np.ndarray) -> np.ndarray:\n    \"\"\"Applies CLAHE to reduce scanner-bias shortcut learning.\"\"\"\n    if img_array.dtype != np.uint8:\n        img_array = (\n            (img_array - img_array.min()) /\n            (img_array.max() - img_array.min()) * 255\n        ).astype(np.uint8)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    return clahe.apply(img_array)\n\n\ndef custom_load_image(abs_path: str, cache_dir: Path) -> Image.Image:\n    \"\"\"Handles DICOM caching, standard loading, and CLAHE preprocessing.\"\"\"\n    if abs_path.endswith('.dcm'):\n        file_name  = os.path.basename(abs_path).replace('.dcm', '.png')\n        cache_path = cache_dir / file_name\n        if cache_path.exists():\n            img_array = np.array(Image.open(cache_path).convert('L'))\n        else:\n            dcm       = pydicom.dcmread(abs_path)\n            img_array = dcm.pixel_array\n            img_array = (\n                (img_array - img_array.min()) /\n                (img_array.max() - img_array.min()) * 255\n            ).astype(np.uint8)\n            Image.fromarray(img_array).save(cache_path)\n    else:\n        img_array = np.array(Image.open(abs_path).convert('L'))\n    return Image.fromarray(apply_clahe(img_array))\n\n\nprint('Image loading utilities defined.')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:59:24.545200Z","iopub.execute_input":"2026-06-22T06:59:24.546158Z","iopub.status.idle":"2026-06-22T06:59:24.553705Z","shell.execute_reply.started":"2026-06-22T06:59:24.546119Z","shell.execute_reply":"2026-06-22T06:59:24.552759Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 6 -- Dataset Class & DataLoader\n# Each sample returns the image + its pre-computed fused SSL\n# conditioning vector looked up from cond_train by row index.\n# -------------------------------------------------------------------\nclass CXRSSLDataset(Dataset):\n    \"\"\"\n    Returns image + fused SSL conditioning vector.\n    cond_vec : (COND_DIM,) = ssl_128d + disease_onehot + severity_onehot + demo_64d\n    disease  : integer label (for discriminator cls_head)\n    \"\"\"\n    def __init__(self, dataframe, cond_array, cache_dir, transform=None):\n        self.df        = dataframe.reset_index(drop=True)\n        self.cond      = cond_array          # numpy (N, COND_DIM)\n        self.cache_dir = cache_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        try:\n            row      = self.df.iloc[idx]\n            image    = custom_load_image(row['abs_path'], self.cache_dir)\n            if self.transform:\n                image = self.transform(image)\n            d_str    = str(row.get('disease', 'No Finding'))\n            disease  = DISEASE_MAP.get(d_str, 0)\n            cond_vec = torch.from_numpy(self.cond[idx])   # (COND_DIM,)\n            return {'image': image, 'disease': disease, 'cond_vec': cond_vec}\n        except Exception:\n            return self.__getitem__(random.randint(0, len(self.df) - 1))\n\n\ngan_transform = transforms.Compose([\n    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5], std=[0.5])   # [0,1] -> [-1,1]\n])\n\ntrain_dataset = CXRSSLDataset(df_train, cond_train, CACHE_DIR, transform=gan_transform)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n    pin_memory=True,\n    drop_last=True\n)\n\nprint(f'DataLoader ready: {len(train_dataset):,} samples, {len(train_loader)} batches/epoch.')\nprint(f'Conditioning vector dim per sample: {COND_DIM}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:59:25.513795Z","iopub.execute_input":"2026-06-22T06:59:25.514219Z","iopub.status.idle":"2026-06-22T06:59:25.543013Z","shell.execute_reply.started":"2026-06-22T06:59:25.514162Z","shell.execute_reply":"2026-06-22T06:59:25.542341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 7 -- Sanity Check: Batch Shape & Conditioning Sub-vectors\n# -------------------------------------------------------------------\nprint('--- DIAGNOSTICS ---')\nbatch = next(iter(train_loader))\nimgs  = batch['image']\n\nprint(f\"Image Batch Shape  : {imgs.shape}\")\nprint(f\"Expected           : [{BATCH_SIZE}, 1, {IMAGE_SIZE}, {IMAGE_SIZE}]\")\nprint(f\"Cond Vec Shape     : {batch['cond_vec'].shape}  | Expected: [{BATCH_SIZE}, {COND_DIM}]\")\nprint(f\"Disease labels     : {batch['disease'][:8].tolist()}\")\n\ncv    = batch['cond_vec']\nd_end = SSL_EMB_DIM + DISEASE_ONEHOT_DIM\ns_end = d_end + SEVERITY_ONEHOT_DIM\nprint('\\nSub-vector stats (first sample):')\nprint(f'  SSL 128d     [0:{SSL_EMB_DIM}]         '\n      f'min={cv[0,:SSL_EMB_DIM].min():.3f}  max={cv[0,:SSL_EMB_DIM].max():.3f}')\nprint(f'  Disease OH   [{SSL_EMB_DIM}:{d_end}]  '\n      f'sum={cv[0,SSL_EMB_DIM:d_end].sum():.1f}  (should be 1.0)')\nprint(f'  Severity OH  [{d_end}:{s_end}]  '\n      f'sum={cv[0,d_end:s_end].sum():.1f}  (should be 1.0)')\nprint(f'  Demo 64d     [{s_end}:{COND_DIM}]  '\n      f'min={cv[0,s_end:].min():.3f}  max={cv[0,s_end:].max():.3f}')\n\ngrid = vutils.make_grid(imgs[:16], nrow=4, padding=2, normalize=True, value_range=(-1, 1))\nplt.figure(figsize=(10, 10))\nplt.axis('off')\nplt.title('Sanity Check: CLAHE + Normalised (first 16 images)')\nplt.imshow(grid[0].cpu().numpy(), cmap='gray')\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:59:28.808551Z","iopub.execute_input":"2026-06-22T06:59:28.809237Z","iopub.status.idle":"2026-06-22T06:59:37.180995Z","shell.execute_reply.started":"2026-06-22T06:59:28.809205Z","shell.execute_reply":"2026-06-22T06:59:37.179665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 8 -- Generator Architecture\n# Identical upsampling structure to original; only input dim changes:\n#   z (128) + fused SSL cond (COND_DIM) -> Linear -> upsample to 256x256\n# -------------------------------------------------------------------\nclass Generator(nn.Module):\n    \"\"\"\n    Input : z (B, Z_DIM) + fused SSL cond (B, COND_DIM)\n    Output: grayscale image (B, 1, 256, 256), range [-1, 1]\n    Upsample path: 16x16 -> 32 -> 64 -> 128 -> 256\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.init_size = IMAGE_SIZE // 16   # 16\n        self.l1 = nn.Sequential(\n            nn.Linear(Z_DIM + COND_DIM, 512 * self.init_size ** 2)\n        )\n        self.conv_blocks = nn.Sequential(\n            nn.BatchNorm2d(512),\n            nn.Upsample(scale_factor=2),                       # 16 -> 32\n            nn.Conv2d(512, 256, 3, stride=1, padding=1),\n            nn.BatchNorm2d(256, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Upsample(scale_factor=2),                       # 32 -> 64\n            nn.Conv2d(256, 128, 3, stride=1, padding=1),\n            nn.BatchNorm2d(128, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Upsample(scale_factor=2),                       # 64 -> 128\n            nn.Conv2d(128, 64, 3, stride=1, padding=1),\n            nn.BatchNorm2d(64, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Upsample(scale_factor=2),                       # 128 -> 256\n            nn.Conv2d(64, CHANNELS, 3, stride=1, padding=1),\n            nn.Tanh()\n        )\n\n    def forward(self, z, cond=None):\n        if cond is not None:\n            z = torch.cat([z, cond], dim=1)\n        out = self.l1(z)\n        out = out.view(out.shape[0], 512, self.init_size, self.init_size)\n        return self.conv_blocks(out)\n\n\n_g = Generator()\nprint(f'Generator: {sum(p.numel() for p in _g.parameters()):,} parameters')\nprint(f'  Input dim: Z_DIM({Z_DIM}) + COND_DIM({COND_DIM}) = {Z_DIM + COND_DIM}')\ndel _g\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T06:59:48.096849Z","iopub.execute_input":"2026-06-22T06:59:48.097469Z","iopub.status.idle":"2026-06-22T06:59:48.429341Z","shell.execute_reply.started":"2026-06-22T06:59:48.097431Z","shell.execute_reply":"2026-06-22T06:59:48.428592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 9 -- Discriminator Architecture\n# Two-head AC-GAN discriminator (identical to original):\n#   adv_head -> 1 scalar  (real / fake)\n#   cls_head -> disease logits  (auxiliary classification)\n# The discriminator sees only the image, not the conditioning vector.\n# -------------------------------------------------------------------\nclass Discriminator(nn.Module):\n    def __init__(self, num_disease_classes: int = 2,\n                 num_severity_classes: int = 4,\n                 demo_emb_dim: int = 64):\n        super().__init__()\n\n        def disc_block(in_f, out_f, bn=True):\n            block = [\n                nn.Conv2d(in_f, out_f, 3, 2, 1),\n                nn.LeakyReLU(0.2, inplace=True),\n                nn.Dropout2d(0.25)\n            ]\n            if bn:\n                block.append(nn.BatchNorm2d(out_f, 0.8))\n            return block\n\n        self.backbone = nn.Sequential(\n            *disc_block(CHANNELS, 64,  bn=False),\n            *disc_block(64,   128),\n            *disc_block(128,  256),\n            *disc_block(256,  512),\n            *disc_block(512, 1024),\n        )\n\n        ds_size = IMAGE_SIZE // (2 ** 5)   # 8\n        flat    = 1024 * ds_size ** 2\n\n        self.adv_head      = nn.Linear(flat, 1)\n        self.cls_head      = nn.Linear(flat, num_disease_classes)\n        self.sev_head      = nn.Linear(flat, num_severity_classes)   # NEW\n        self.demo_head     = nn.Linear(flat, demo_emb_dim)           # NEW\n\n    def forward(self, x, cond=None):\n        out = self.backbone(x).view(x.shape[0], -1)\n        return (\n            self.adv_head(out),\n            self.cls_head(out),\n            self.sev_head(out),    # NEW\n            self.demo_head(out)    # NEW\n        )\n\n\n_d = Discriminator(\n    num_disease_classes=DISEASE_ONEHOT_DIM,\n    num_severity_classes=SEVERITY_ONEHOT_DIM,\n    demo_emb_dim=DEMO_EMB_DIM\n)\nprint(f'Discriminator: {sum(p.numel() for p in _d.parameters()):,} parameters')\nprint(f'  adv_head  -> 1')\nprint(f'  cls_head  -> {DISEASE_ONEHOT_DIM} (disease)')\nprint(f'  sev_head  -> {SEVERITY_ONEHOT_DIM} (severity)')\nprint(f'  demo_head -> {DEMO_EMB_DIM} (demographics)')\ndel _d\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T07:00:53.021683Z","iopub.execute_input":"2026-06-22T07:00:53.021983Z","iopub.status.idle":"2026-06-22T07:00:53.110048Z","shell.execute_reply.started":"2026-06-22T07:00:53.021958Z","shell.execute_reply":"2026-06-22T07:00:53.109183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 10 -- Instantiate Models, Optimizers & Multi-GPU Wrapping\n# -------------------------------------------------------------------\ngenerator     = Generator().to(device)\ndiscriminator = Discriminator(\n    num_disease_classes=DISEASE_ONEHOT_DIM,\n    num_severity_classes=SEVERITY_ONEHOT_DIM,\n    demo_emb_dim=DEMO_EMB_DIM\n).to(device)\n\nif torch.cuda.device_count() > 1:\n    print(f'{torch.cuda.device_count()} GPUs detected -- wrapping with DataParallel')\n    generator     = nn.DataParallel(generator)\n    discriminator = nn.DataParallel(discriminator)\n\noptimizer_G = torch.optim.Adam(generator.parameters(),     lr=2e-4, betas=(0.5, 0.999))\noptimizer_D = torch.optim.Adam(discriminator.parameters(), lr=1e-4, betas=(0.5, 0.999))\n\nprint('Models and optimizers ready.')\nprint(f'  Generator     params: {sum(p.numel() for p in generator.parameters()):,}')\nprint(f'  Discriminator params: {sum(p.numel() for p in discriminator.parameters()):,}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T07:01:05.885001Z","iopub.execute_input":"2026-06-22T07:01:05.885739Z","iopub.status.idle":"2026-06-22T07:01:06.327333Z","shell.execute_reply.started":"2026-06-22T07:01:05.885705Z","shell.execute_reply":"2026-06-22T07:01:06.326646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 11 -- DINOv2 Perceptual Loss Extractor (Frozen)\n# Used only as a training regulariser (loss_G_ssl), NOT as cond input.\n# Penalises the generator if fake image DINOv2 features drift from\n# the corresponding real image features.\n# -------------------------------------------------------------------\nprint('Loading frozen DINOv2 (ViT-S/14) for perceptual loss...')\n\ndino_model = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14')\ndino_model = dino_model.to(device)\n\nfor param in dino_model.parameters():\n    param.requires_grad = False\ndino_model.eval()\n\nif torch.cuda.device_count() > 1:\n    dino_model = nn.DataParallel(dino_model)\n\n\ndef extract_ssl_features(imgs: torch.Tensor, model) -> torch.Tensor:\n    \"\"\"Resize grayscale -> RGB 112x112 (half size to save VRAM), extract DINOv2 CLS token.\"\"\"\n    imgs_resized = F.interpolate(imgs, size=(112, 112), mode='bilinear', align_corners=False)\n    imgs_rgb     = imgs_resized.repeat(1, 3, 1, 1)   # (B,1,H,W) -> (B,3,H,W)\n    return model(imgs_rgb)\n\n\nprint('DINOv2 perceptual loss extractor ready (frozen).')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T07:01:13.498412Z","iopub.execute_input":"2026-06-22T07:01:13.499132Z","iopub.status.idle":"2026-06-22T07:01:14.689415Z","shell.execute_reply.started":"2026-06-22T07:01:13.499098Z","shell.execute_reply":"2026-06-22T07:01:14.688665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 12 -- Fixed Conditioning Vectors for Visual Monitoring\n# We sample one real conditioning vector per disease class and tile\n# it across 4 noise draws so the grid shows different noise seeds\n# under the same SSL condition.\n#\n# Grid layout (nrow=4):\n#   Row 1 (idx  0-3) : No Finding\n#   Row 2 (idx  4-7) : Pneumonia\n#   Row 3 (idx  8-11): COVID-19\n#   Row 4 (idx 12-15): No Finding  (repeat for visual symmetry)\n# -------------------------------------------------------------------\nfixed_z = torch.randn(16, Z_DIM, device=device)\n\nfixed_cond_np = np.zeros((16, COND_DIM), dtype=np.float32)\n\nfor class_idx, class_name in enumerate(['No Finding', 'Pneumonia', 'COVID-19']):\n    mask          = df_train['disease'].map(DISEASE_MAP) == class_idx\n    class_indices = np.where(mask)[0]\n    if len(class_indices) == 0:\n        print(f'WARNING: No samples for class {class_name}, using zeros.')\n        continue\n    rep_idx = class_indices[len(class_indices) // 2]   # median sample\n    rep_vec = cond_train[rep_idx]                       # (COND_DIM,)\n    row_start = class_idx * 4\n    fixed_cond_np[row_start:row_start + 4] = rep_vec\n\nfixed_cond_np[12:16] = fixed_cond_np[0]   # Row 4 = No Finding repeat\n\nfixed_cond = torch.from_numpy(fixed_cond_np).to(device)\n\nprint('Fixed noise and SSL conditioning vectors set.')\nprint(f'  fixed_z    : {fixed_z.shape}')\nprint(f'  fixed_cond : {fixed_cond.shape}')\nprint('Grid layout (nrow=4):  No Finding | Pneumonia | COVID-19 | No Finding')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-22T07:01:17.444914Z","iopub.execute_input":"2026-06-22T07:01:17.445533Z","iopub.status.idle":"2026-06-22T07:01:17.481460Z","shell.execute_reply.started":"2026-06-22T07:01:17.445499Z","shell.execute_reply":"2026-06-22T07:01:17.480630Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 13 -- Training Loop\n#\n# Key differences from the original disease-only GAN:\n#   - cond_vec comes from npy-backed CXRSSLDataset (SSL embeddings)\n#     instead of a 3-d one-hot disease label\n#   - fake conditioning is drawn by randomly sampling rows from the\n#     full cond_train tensor, preserving the joint SSL distribution\n#   - Disease integer labels are derived from the disease one-hot\n#     sub-vector of the sampled fake cond, so cls_head training is\n#     still anchored to the correct class\n#   - DINOv2 perceptual regularisation (LAMBDA_SSL) is unchanged\n# -------------------------------------------------------------------\nEPOCHS      = 150\nSTART_EPOCH = 0\n\ncriterion_cls  = nn.CrossEntropyLoss().to(device)\ncriterion_sev  = nn.CrossEntropyLoss().to(device)\ncriterion_demo = nn.MSELoss().to(device)\n\n# ── Optional checkpoint resume ────────────────────────────────────\nCHECKPOINT_PATH = None\n\ndef load_state_dict_flexible(model, state_dict):\n    target    = model.module if isinstance(model, nn.DataParallel) else model\n    first_key = next(iter(state_dict))\n    if first_key.startswith('module.'):\n        state_dict = {k[len('module.'):]: v for k, v in state_dict.items()}\n    target.load_state_dict(state_dict)\n\nif CHECKPOINT_PATH and Path(CHECKPOINT_PATH).exists():\n    ckpt = torch.load(CHECKPOINT_PATH, map_location=device)\n    load_state_dict_flexible(generator,     ckpt['generator_state_dict'])\n    load_state_dict_flexible(discriminator, ckpt['discriminator_state_dict'])\n    optimizer_G.load_state_dict(ckpt['optimizer_G_state_dict'])\n    optimizer_D.load_state_dict(ckpt['optimizer_D_state_dict'])\n    START_EPOCH = ckpt['epoch']\n    print(f'Resumed from epoch {START_EPOCH}. Training to epoch {EPOCHS}.')\nelse:\n    print(f'Training from scratch -- epoch 1 to {EPOCHS}.')\n\n# Keep cond matrix on CPU -- sample rows and move to GPU per batch\nall_cond_np = cond_train   # numpy (N, COND_DIM)\nN_COND      = all_cond_np.shape[0]\n\n# Sub-vector slice boundaries\nD_START = SSL_EMB_DIM\nD_END   = SSL_EMB_DIM + DISEASE_ONEHOT_DIM\nS_END   = D_END + SEVERITY_ONEHOT_DIM\n# demo  = cond[:, S_END:COND_DIM]  shape (B, 64)\n\nfor epoch in range(START_EPOCH, EPOCHS):\n    generator.train()\n    discriminator.train()\n    acc_d_list = []\n\n    for i, batch in enumerate(tqdm(train_loader, desc=f'Epoch {epoch+1}/{EPOCHS}', leave=False)):\n        try:\n            real_imgs    = batch['image'].to(device)\n            real_disease = batch['disease'].to(device)\n            real_cond    = batch['cond_vec'].to(device)    # (B, COND_DIM)\n\n            real_severity = torch.argmax(real_cond[:, D_END:S_END], dim=1)\n            real_demo     = real_cond[:, S_END:]\n\n            B = real_imgs.size(0)\n\n            # ── Train Discriminator ──────────────────────────────────\n            optimizer_D.zero_grad()\n\n            real_adv, real_cls, real_sev, real_demo_pred = discriminator(real_imgs)\n\n            loss_D_real_adv  = F.relu(0.9 - real_adv).mean()\n            loss_D_real_cls  = criterion_cls(real_cls, real_disease)\n            loss_D_real_sev  = criterion_sev(real_sev, real_severity)\n            loss_D_real_demo = criterion_demo(real_demo_pred, real_demo)\n\n            # Sample fake cond on CPU, move to GPU\n            rand_idx      = np.random.randint(0, N_COND, size=B)\n            fake_cond     = torch.from_numpy(all_cond_np[rand_idx]).to(device)\n            fake_disease  = torch.argmax(fake_cond[:, D_START:D_END], dim=1)\n            fake_severity = torch.argmax(fake_cond[:, D_END:S_END],   dim=1)\n            fake_demo     = fake_cond[:, S_END:]\n\n            z         = torch.randn(B, Z_DIM, device=device)\n            fake_imgs = generator(z, fake_cond)\n\n            fake_adv, _, _, _ = discriminator(fake_imgs.detach())\n            loss_D_fake_adv   = F.relu(0.9 + fake_adv).mean()\n\n            loss_D_adv = loss_D_real_adv + loss_D_fake_adv\n            loss_D = (loss_D_adv\n                      + LAMBDA_CLS  * loss_D_real_cls\n                      + LAMBDA_SEV  * loss_D_real_sev\n                      + LAMBDA_DEMO * loss_D_real_demo)\n            loss_D.backward()\n            optimizer_D.step()\n\n            # ── Train Generator ──────────────────────────────────────\n            optimizer_G.zero_grad()\n\n            gen_adv, gen_cls, gen_sev, gen_demo = discriminator(fake_imgs)\n\n            loss_G_adv  = -gen_adv.mean()\n            loss_G_cls  = criterion_cls(gen_cls, fake_disease)\n            loss_G_sev  = criterion_sev(gen_sev, fake_severity)\n            loss_G_demo = criterion_demo(gen_demo, fake_demo)\n\n            # DINOv2 perceptual loss -- only every SSL_EVERY batches\n            if i % SSL_EVERY == 0:\n                with torch.no_grad():\n                    real_features = extract_ssl_features(real_imgs, dino_model)\n                fake_features = extract_ssl_features(fake_imgs, dino_model)\n                loss_G_ssl    = F.l1_loss(fake_features, real_features)\n            else:\n                loss_G_ssl = torch.tensor(0.0, device=device)\n\n            loss_G = (loss_G_adv\n                      + LAMBDA_CLS  * loss_G_cls\n                      + LAMBDA_SEV  * loss_G_sev\n                      + LAMBDA_DEMO * loss_G_demo\n                      + LAMBDA_SSL  * loss_G_ssl)\n            loss_G.backward()\n            optimizer_G.step()\n\n            acc_d_list.append(\n                (torch.argmax(gen_cls, dim=1) == fake_disease).float().mean().item()\n            )\n\n            # Periodic VRAM cleanup inside the batch loop\n            if i % 20 == 0:\n                torch.cuda.empty_cache()\n\n        except Exception as e:\n            print(f'Skipped corrupted batch {i}: {e}')\n            continue\n\n    if acc_d_list:\n        avg_acc = sum(acc_d_list) / len(acc_d_list) * 100\n        print(\n            f'[Ep {epoch+1:03d}/{EPOCHS}] '\n            f'D: {loss_D_adv.item():.3f}  '\n            f'G_adv: {loss_G_adv.item():.3f}  '\n            f'G_cls: {loss_G_cls.item():.3f}  '\n            f'G_sev: {loss_G_sev.item():.3f}  '\n            f'G_demo: {loss_G_demo.item():.4f}  '\n            f'SSL: {loss_G_ssl.item():.4f}  '\n            f'Dis Cls Acc: {avg_acc:.1f}%'\n        )\n\n    if (epoch + 1) % 25 == 0:\n        ckpt_out = {\n            'epoch'                    : epoch + 1,\n            'generator_state_dict'     : generator.state_dict(),\n            'discriminator_state_dict' : discriminator.state_dict(),\n            'optimizer_G_state_dict'   : optimizer_G.state_dict(),\n            'optimizer_D_state_dict'   : optimizer_D.state_dict(),\n            'loss_G'                   : loss_G.item(),\n            'loss_D'                   : loss_D.item(),\n            'cond_dim'                 : COND_DIM,\n        }\n        ckpt_path = f'/kaggle/working/ckpt_ssl_gan_ep{epoch+1}.pth'\n        torch.save(ckpt_out, ckpt_path)\n        print(f'Checkpoint saved -> {ckpt_path}')\n\n    if epoch == START_EPOCH or (epoch + 1) % 5 == 0:\n        generator.eval()\n        with torch.no_grad():\n            sample_imgs = generator(fixed_z, fixed_cond)\n            grid = vutils.make_grid(\n                sample_imgs, nrow=4, padding=2,\n                normalize=True, value_range=(-1, 1)\n            )\n        from IPython.display import display\n        plt.figure(figsize=(8, 8))\n        plt.axis('off')\n        plt.title(\n            f'Epoch {epoch+1}  |  Row1: No Finding  Row2: Pneumonia  Row3: No Finding  Row4: Pneumonia',\n            fontsize=9\n        )\n        plt.imshow(grid[0].cpu().numpy(), cmap='gray')\n        save_path = OUTPUT_GRID_DIR / f'grid_ep{epoch+1:04d}.png'\n        plt.savefig(save_path, bbox_inches='tight', dpi=100)\n        display(plt.gcf())\n        plt.close()\n        generator.train()\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint('\\nTraining complete.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 14 -- Final Sample Grid per Disease Class\n# 8 samples per class using the representative fixed cond vector.\n# -------------------------------------------------------------------\ngenerator.eval()\n\nfig, axes = plt.subplots(3, 8, figsize=(20, 8))\nfig.suptitle('Final Generated Samples -- 8 per Disease Class', fontsize=14, y=1.01)\n\nwith torch.no_grad():\n    for class_idx, class_name in enumerate(CLASS_NAMES):\n        rep_cond = fixed_cond[class_idx * 4].unsqueeze(0).expand(8, -1)  # (8, COND_DIM)\n        z        = torch.randn(8, Z_DIM, device=device)\n        imgs     = generator(z, rep_cond)\n        for j in range(8):\n            ax     = axes[class_idx][j]\n            img_np = (imgs[j, 0].cpu().numpy() + 1) / 2   # [-1,1] -> [0,1]\n            ax.imshow(img_np, cmap='gray', vmin=0, vmax=1)\n            ax.axis('off')\n            if j == 0:\n                ax.set_title(class_name, fontsize=11, fontweight='bold', loc='left')\n\nplt.tight_layout()\nfinal_path = OUTPUT_GRID_DIR / 'final_samples_per_class.png'\nplt.savefig(final_path, bbox_inches='tight', dpi=120)\nplt.show()\nprint(f'Final sample grid saved -> {final_path}')\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------------\n# Cell 15 -- Save Final Model Weights\n# -------------------------------------------------------------------\nfinal_ckpt = {\n    'generator_state_dict'     : generator.state_dict(),\n    'discriminator_state_dict' : discriminator.state_dict(),\n    'optimizer_G_state_dict'   : optimizer_G.state_dict(),\n    'optimizer_D_state_dict'   : optimizer_D.state_dict(),\n    'z_dim'                    : Z_DIM,\n    'cond_dim'                 : COND_DIM,\n    'ssl_emb_dim'              : SSL_EMB_DIM,\n    'disease_onehot_dim'       : DISEASE_ONEHOT_DIM,\n    'severity_onehot_dim'      : SEVERITY_ONEHOT_DIM,\n    'demo_emb_dim'             : DEMO_EMB_DIM,\n    'image_size'               : IMAGE_SIZE,\n    'class_names'              : CLASS_NAMES,\n    'disease_map'              : DISEASE_MAP,\n}\n\nfinal_path = '/kaggle/working/gan_ssl_conditioned_final.pth'\ntorch.save(final_ckpt, final_path)\nprint(f'Final weights saved -> {final_path}')\nprint(f'  COND_DIM breakdown: '\n      f'SSL({SSL_EMB_DIM}) + Disease({DISEASE_ONEHOT_DIM}) '\n      f'+ Severity({SEVERITY_ONEHOT_DIM}) + Demo({DEMO_EMB_DIM}) = {COND_DIM}')\n","metadata":{},"outputs":[],"execution_count":null}]}