{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042},{"sourceType":"datasetVersion","sourceId":23812,"datasetId":17810,"databundleVersionId":23851},{"sourceType":"datasetVersion","sourceId":2169393,"datasetId":1302315,"databundleVersionId":2210641},{"sourceType":"datasetVersion","sourceId":6717213,"datasetId":1317048,"databundleVersionId":6801677},{"sourceType":"datasetVersion","sourceId":15762157,"datasetId":10101314,"databundleVersionId":16706329},{"sourceType":"datasetVersion","sourceId":18613,"datasetId":5839,"databundleVersionId":18613},{"sourceType":"datasetVersion","sourceId":15735855,"datasetId":10082918,"databundleVersionId":16677853}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport time\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport torchvision.utils as vutils\nimport pandas as pd\nimport numpy as np\nimport pyarrow.parquet as pq\nfrom PIL import Image\nimport cv2\nimport pydicom\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom pathlib import Path\n\n# --- HARDCODED ARCHITECTURE CONSTANTS ---\nIMAGE_SIZE = 256  \nCHANNELS = 1\nZ_DIM = 100\nCOND_DIM = 199    # UPGRADED: The 199-dim embedding vector\nBATCH_SIZE = 64   \nNUM_WORKERS = 4\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"✅ Using device: {device}\")\n\n# --- YOUR EXACT PATH MAPPINGS ---\nPARQUET_PATH = \"/kaggle/input/datasets/spacecypher/finaldata/cgan_input_full.parquet\"\nCACHE_DIR = Path('/kaggle/working/dcm_cache')\nCACHE_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}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:02:54.243783Z","iopub.execute_input":"2026-04-16T19:02:54.244141Z","iopub.status.idle":"2026-04-16T19:03:07.285845Z","shell.execute_reply.started":"2026-04-16T19:02:54.244082Z","shell.execute_reply":"2026-04-16T19:03:07.285175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\n\ndef resolve_path(row):\n    \"\"\"Maps a parquet-relative path to its absolute Kaggle dataset path using strings.\"\"\"\n    ds_name = row.get(\"dataset\") \n    root = DATASET_ROOTS.get(ds_name)\n    \n    if root is None:\n        return str(row[\"path\"])\n        \n    p = str(row[\"path\"]).lstrip('/') \n    \n    # RSNA specific logic (append .dcm if missing)\n    if ds_name == \"RSNA\" and not p.endswith('.dcm'):\n        return str(Path(root) / (Path(p).stem + \".dcm\"))\n        \n    return str(Path(root) / p)\n\ndef load_and_sample_parquet(parquet_path, max_per_class=1000):\n    print(f\"Inspecting schema for {parquet_path}...\")\n    \n    # Load via pandas\n    df_raw = pd.read_parquet(parquet_path)\n    \n    # --- THE FIX: Check for the columns that actually exist in this file ---\n    df_raw = df_raw.dropna(subset=['disease', 'severity_label', 'cgan_conditioning']).reset_index(drop=True)\n    \n    # Resolve absolute paths\n    df_raw['abs_path'] = df_raw.apply(resolve_path, axis=1)\n    \n    print(f\"Applying stratified sampling (max {max_per_class} per disease class)...\")\n    samples = []\n    for _, group in df_raw.groupby('disease'):\n        samples.append(group.sample(n=min(len(group), max_per_class), random_state=42))\n    df_sample = pd.concat(samples).reset_index(drop=True)\n    \n    # Verification: Check Paths\n    print(\"Verifying file existence on disk...\")\n    valid_rows = []\n    missing_paths_preview = []\n    \n    for _, row in tqdm(df_sample.iterrows(), total=len(df_sample), desc=\"Checking paths\"):\n        if Path(row['abs_path']).exists():\n            valid_rows.append(row)\n        elif len(missing_paths_preview) < 5:\n            missing_paths_preview.append((row.get('dataset', 'Unknown'), row['path'], row['abs_path']))\n            \n    df_verified = pd.DataFrame(valid_rows)\n    \n    if len(df_verified) == 0:\n        print(\"\\n❌ CRITICAL ERROR: 0 files found! Path resolution failed.\")\n        print(\"Here are the first 5 paths it tried to find:\\n\")\n        for ds_name, orig, attempt in missing_paths_preview:\n            print(f\"Dataset '{ds_name}' Attempt: {attempt}\")\n    else:\n        print(f\"\\n✅ Final verified dataset size: {len(df_verified)} samples\")\n        \n        # ==========================================\n        # --- THE FIX: Updated Distribution Printouts ---\n        # ==========================================\n        print(\"\\n--- FINAL DATASET DISTRIBUTION ---\")\n        \n        print(\"\\n1. Disease Distribution (Target: 3 Classes):\")\n        print(df_verified['disease'].value_counts())\n        \n        print(\"\\n2. Severity Distribution:\")\n        print(df_verified['severity_label'].value_counts())\n        print(\"----------------------------------\\n\")\n        \n    return df_verified\n\n# Execute loading and sampling\ntrain_df = load_and_sample_parquet(PARQUET_PATH, max_per_class=5000)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:03:07.287229Z","iopub.execute_input":"2026-04-16T19:03:07.287621Z","iopub.status.idle":"2026-04-16T19:03:55.709418Z","shell.execute_reply.started":"2026-04-16T19:03:07.287596Z","shell.execute_reply":"2026-04-16T19:03:55.708577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_clahe(img_array):\n    \"\"\"Applies CLAHE to fix scanner bias (shortcut learning).\"\"\"\n    if img_array.dtype != np.uint8:\n        img_array = ((img_array - img_array.min()) / (img_array.max() - img_array.min()) * 255).astype(np.uint8)\n    \n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    return clahe.apply(img_array)\n\ndef custom_load_image(abs_path, cache_dir):\n    \"\"\"Handles DICOM caching, standard loading, and CLAHE.\"\"\"\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        \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 = ((img_array - img_array.min()) / (img_array.max() - img_array.min()) * 255).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        \n    clahe_img = apply_clahe(img_array)\n    return Image.fromarray(clahe_img)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:03:55.710387Z","iopub.execute_input":"2026-04-16T19:03:55.710669Z","iopub.status.idle":"2026-04-16T19:03:55.717824Z","shell.execute_reply.started":"2026-04-16T19:03:55.710646Z","shell.execute_reply":"2026-04-16T19:03:55.717324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.utils as vutils\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\n\n# --- LABEL MAPPINGS ---\nDISEASE_MAP = {\"No Finding\": 0, \"Normal\": 0, \"Pneumonia\": 1, \"COVID-19\": 2}\n\ndef get_severity_idx(label_str):\n    label_str = str(label_str).lower()\n    if 'mild' in label_str: return 1\n    if 'moderate' in label_str: return 2\n    if 'severe' in label_str: return 3\n    return 0 # Normal\n\nclass AdvancedCXRConditionalDataset(Dataset):\n    def __init__(self, dataframe, cache_dir, transform=None):\n        self.df = dataframe\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            \n            # --- THE CRITICAL FIX: Use 'abs_path' instead of 'path' ---\n            # abs_path contains the full Kaggle directory route generated in Block 2\n            image = custom_load_image(row['abs_path'], self.cache_dir) \n\n            if self.transform:\n                image = self.transform(image)\n\n            # Map Discrete Classes for the Discriminator\n            disease = DISEASE_MAP.get(str(row.get('disease', 'No Finding')), 0)\n            severity = get_severity_idx(row.get('severity_label', 'Normal'))\n            \n            # Extract the 199-dim vector\n            cond_vec = np.array(row['cgan_conditioning'], dtype=np.float32)\n            cond_vec = torch.tensor(cond_vec)\n\n            return {\n                \"image\": image,\n                \"disease\": disease,\n                \"severity\": severity,\n                \"cond_vec\": cond_vec\n            }\n            \n        except Exception as e:\n            # If the image is missing, truncated, or corrupted, grab a random new one!\n            random_idx = random.randint(0, len(self.df) - 1)\n            return self.__getitem__(random_idx)\n\n# --- DATALOADER SETUP ---\ngan_transform = transforms.Compose([\n    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5], std=[0.5]) \n])\n\ntrain_dataset = AdvancedCXRConditionalDataset(train_df, 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\n# --- RUN DIAGNOSTICS & PLOT GRID ---\nprint(\"\\n--- RUNNING STAGE 0 DIAGNOSTICS ---\")\nbatch = next(iter(train_loader))\nimgs = batch[\"image\"]\n\nprint(f\"Image Batch Shape: {imgs.shape} | Expected: [{BATCH_SIZE}, 1, {IMAGE_SIZE}, {IMAGE_SIZE}]\")\nprint(f\"Cond Vec Shape:    {batch['cond_vec'].shape}  | Expected: [{BATCH_SIZE}, 199]\")\n\nplt.figure(figsize=(12, 12))\nplt.axis(\"off\")\nplt.title(\"Sanity Check: CLAHE Applied & Normalized\")\ngrid = vutils.make_grid(imgs[:16], nrow=4, padding=2, normalize=True, value_range=(-1, 1))\nplt.imshow(grid[0].cpu().numpy(), cmap='gray')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:03:55.719349Z","iopub.execute_input":"2026-04-16T19:03:55.719754Z","iopub.status.idle":"2026-04-16T19:04:05.942348Z","shell.execute_reply.started":"2026-04-16T19:03:55.719715Z","shell.execute_reply":"2026-04-16T19:04:05.941514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\n# --- UPDATED ARCHITECTURE CONSTANTS ---\nZ_DIM = 100       # Ensure this matches whatever Z_DIM you set in Block 1\nCOND_DIM = 199    # THE GOLDEN TICKET (128d deep features + 64d demographics + 7d labels)\n\nclass Generator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # For 256 resolution: 256 // 16 = 16\n        self.init_size = IMAGE_SIZE // 16  \n        \n        # Input: Z (100) + Cond (199) = 299\n        self.l1 = nn.Sequential(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), # 16x16 -> 32x32\n            nn.Conv2d(512, 256, 3, stride=1, padding=1),\n            nn.BatchNorm2d(256, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            \n            nn.Upsample(scale_factor=2), # 32x32 -> 64x64\n            nn.Conv2d(256, 128, 3, stride=1, padding=1),\n            nn.BatchNorm2d(128, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            \n            nn.Upsample(scale_factor=2), # 64x64 -> 128x128\n            nn.Conv2d(128, 64, 3, stride=1, padding=1),\n            nn.BatchNorm2d(64, 0.8),\n            nn.LeakyReLU(0.2, inplace=True),\n            \n            nn.Upsample(scale_factor=2), # 128x128 -> 256x256\n            nn.Conv2d(64, CHANNELS, 3, stride=1, padding=1),\n            nn.Tanh() # Output range [-1, 1] to match the transform normalization\n        )\n\n    def forward(self, z, cond=None):\n        if cond is not None:\n            z = torch.cat([z, cond], dim=1)\n        \n        out = self.l1(z)\n        out = out.view(out.shape[0], 512, self.init_size, self.init_size)\n        img = self.conv_blocks(out)\n        return img\n\nclass Discriminator(nn.Module):\n    # Updated default to 7 (3 Disease + 4 Severity)\n    def __init__(self, num_classes=7): \n        super().__init__()\n\n        def discriminator_block(in_filters, out_filters, bn=True):\n            block = [nn.Conv2d(in_filters, out_filters, 3, 2, 1), \n                     nn.LeakyReLU(0.2, inplace=True), \n                     nn.Dropout2d(0.25)]\n            if bn:\n                block.append(nn.BatchNorm2d(out_filters, 0.8))\n            return block\n\n        self.backbone = nn.Sequential(\n            *discriminator_block(CHANNELS, 64, bn=False), # 256x256 -> 128x128\n            *discriminator_block(64, 128),                # 128x128 -> 64x64\n            *discriminator_block(128, 256),               # 64x64 -> 32x32\n            *discriminator_block(256, 512),               # 32x32 -> 16x16\n            *discriminator_block(512, 1024),              # 16x16 -> 8x8\n        )\n\n        # The downsampled size is 256 / (2^5) = 8\n        ds_size = IMAGE_SIZE // (2 ** 5)\n        \n        # Two heads\n        self.adv_head = nn.Linear(1024 * ds_size ** 2, 1)          # Real vs Fake\n        self.cls_head = nn.Linear(1024 * ds_size ** 2, num_classes) # Classification grading\n\n    def forward(self, x, cond=None):\n        out = self.backbone(x)\n        out = out.view(out.shape[0], -1)\n        adv = self.adv_head(out)\n        cls = self.cls_head(out)\n        return adv, cls\n\nprint(\"✅ Generative Architecture Defined for 199-Dim & 256x256.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:04:05.943787Z","iopub.execute_input":"2026-04-16T19:04:05.944079Z","iopub.status.idle":"2026-04-16T19:04:05.956956Z","shell.execute_reply.started":"2026-04-16T19:04:05.944052Z","shell.execute_reply":"2026-04-16T19:04:05.956371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\n# Initialize Models \ngenerator = Generator().to(device)\ndiscriminator = Discriminator(num_classes=7).to(device)\n\n# --- DUAL GPU UPGRADE ---\nif torch.cuda.device_count() > 1:\n    print(f\"🚀 Dual GPUs Detected! Using {torch.cuda.device_count()} GPUs!\")\n    generator = nn.DataParallel(generator)\n    discriminator = nn.DataParallel(discriminator)\n\n# Optimizers\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\nOUTPUT_GRID_DIR = Path('/kaggle/working/stage5_grids')\nOUTPUT_GRID_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"✅ Stage 5 Models Initialized for Multi-GPU.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:04:05.958066Z","iopub.execute_input":"2026-04-16T19:04:05.958404Z","iopub.status.idle":"2026-04-16T19:04:06.408992Z","shell.execute_reply.started":"2026-04-16T19:04:05.958371Z","shell.execute_reply":"2026-04-16T19:04:06.408183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nimport torch.nn as nn\n\nprint(\"Downloading and initializing DINOv2...\")\n\ndino_model = torch.hub.load('facebookresearch/dinov2', 'dinov2_vits14')\ndino_model = dino_model.to(device)\n\n# Freeze the model\nfor param in dino_model.parameters():\n    param.requires_grad = False\ndino_model.eval()\n\n# --- DUAL GPU UPGRADE ---\nif torch.cuda.device_count() > 1:\n    dino_model = nn.DataParallel(dino_model)\n\ndef extract_ssl_features(imgs, model):\n    imgs_resized = F.interpolate(imgs, size=(224, 224), mode='bilinear', align_corners=False)\n    imgs_rgb = imgs_resized.repeat(1, 3, 1, 1)\n    features = model(imgs_rgb)\n    return features\n\nprint(\"✅ Frozen DINOv2 Multi-GPU Extractor Ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:04:06.410066Z","iopub.execute_input":"2026-04-16T19:04:06.410400Z","iopub.status.idle":"2026-04-16T19:04:08.810818Z","shell.execute_reply.started":"2026-04-16T19:04:06.410365Z","shell.execute_reply":"2026-04-16T19:04:08.810170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nimport torch\nimport torch.nn as nn\nfrom tqdm import tqdm\nimport gc \nimport torchvision.utils as vutils\nimport matplotlib.pyplot as plt\n\nEPOCHS_STAGE4 = 150  \nLAMBDA_CLS = 1.0 \nLAMBDA_SSL = 5.0  \n\nprint(f\"--- STARTING STAGE 5: 199-DIM EMBEDDING SCALE-UP ({EPOCHS_STAGE4} EPOCHS) ---\")\n\ncriterion_cls = nn.CrossEntropyLoss().to(device)\n\n# Setup a fixed visual grid using real 199-dim vectors from the first batch\ngenerator.eval()\nfixed_batch = next(iter(train_loader))\nfixed_z_stage4 = torch.randn(16, Z_DIM, device=device)\nfixed_cond_stage4 = fixed_batch[\"cond_vec\"][:16].to(device) \ngenerator.train()\n\nfor epoch in range(EPOCHS_STAGE4):\n    acc_d, acc_s = [], []\n    \n    for i, batch in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS_STAGE4}\")):\n        try:\n            real_imgs = batch[\"image\"].to(device)\n            real_disease = batch[\"disease\"].to(device) \n            real_severity = batch[\"severity\"].to(device)\n            real_cond_vecs = batch[\"cond_vec\"].to(device)\n            \n            # ==========================================\n            #  Train Discriminator \n            # ==========================================\n            optimizer_D.zero_grad()\n            real_adv, real_cls = discriminator(real_imgs)\n            loss_D_real_adv = F.relu(0.9 - real_adv).mean()\n            \n            # Grade the real images on Disease (0:3) and Severity (3:7)\n            loss_D_real_cls = criterion_cls(real_cls[:, 0:3], real_disease) + \\\n                              criterion_cls(real_cls[:, 3:7], real_severity)\n\n            # Generate Fake Images using SHUFFLED real conditions\n            z = torch.randn(real_imgs.shape[0], Z_DIM, device=device)\n            shuffled_idx = torch.randperm(real_imgs.shape[0], device=device)\n            \n            fake_cond_vecs = real_cond_vecs[shuffled_idx]\n            fake_disease = real_disease[shuffled_idx]\n            fake_severity = real_severity[shuffled_idx]\n\n            fake_imgs = generator(z, fake_cond_vecs)\n            \n            fake_adv, fake_cls = 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 + (LAMBDA_CLS * loss_D_real_cls)\n            loss_D.backward()\n            optimizer_D.step()\n\n            # ==========================================\n            #  Train Generator \n            # ==========================================\n            optimizer_G.zero_grad()\n            gen_adv, gen_cls = discriminator(fake_imgs)\n            loss_G_adv = -gen_adv.mean()\n            \n            loss_G_cls = criterion_cls(gen_cls[:, 0:3], fake_disease) + \\\n                         criterion_cls(gen_cls[:, 3:7], fake_severity)\n\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            \n            loss_G_ssl = F.l1_loss(fake_features, real_features)\n\n            loss_G = loss_G_adv + (LAMBDA_CLS * loss_G_cls) + (LAMBDA_SSL * loss_G_ssl)\n            loss_G.backward()\n            optimizer_G.step()\n            \n            acc_d.append((torch.argmax(gen_cls[:, 0:3], dim=1) == fake_disease).float().mean().item())\n            acc_s.append((torch.argmax(gen_cls[:, 3:7], dim=1) == fake_severity).float().mean().item())\n\n        except Exception as e:\n            print(f\"\\n⚠️ Skipped corrupted batch {i}: {e}\")\n            continue\n\n    # --- Diagnostics ---\n    if len(acc_d) > 0:\n        avg_d, avg_s = sum(acc_d)/len(acc_d)*100, sum(acc_s)/len(acc_s)*100\n        print(f\"\\n[Ep {epoch+1}/{EPOCHS_STAGE4}] [D: {loss_D_adv.item():.2f}] [G: {loss_G_adv.item():.2f}] [SSL: {loss_G_ssl.item():.4f}] | Dis: {avg_d:.1f}% | Sev: {avg_s:.1f}%\")\n    \n    # --- Checkpointing ---\n    if (epoch + 1) % 25 == 0:\n        checkpoint = {\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        }\n        torch.save(checkpoint, f'/kaggle/working/checkpoint_highres_ep{epoch+1}.pth')\n        print(f\"💾 Checkpoint Saved: Epoch {epoch+1}\")\n    \n    # --- Visual Grid ---\n    if epoch == 0 or (epoch + 1) % 5 == 0:\n        generator.eval() \n        with torch.no_grad():\n            sample_imgs = generator(fixed_z_stage4, fixed_cond_stage4)\n            grid = vutils.make_grid(sample_imgs, nrow=4, padding=2, normalize=True, value_range=(-1, 1))\n            plt.figure(figsize=(8, 8))\n            plt.axis(\"off\")\n            plt.title(f\"199-Dim Scale-Up - Epoch {epoch+1}\")\n            plt.imshow(grid[0].cpu().numpy(), cmap='gray')\n            plt.show()\n        generator.train() \n        \n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(\"✅ 199-Dim Scale-Up Complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-16T19:04:08.811813Z","iopub.execute_input":"2026-04-16T19:04:08.812050Z","iopub.status.idle":"2026-04-16T19:13:02.284347Z","shell.execute_reply.started":"2026-04-16T19:04:08.812028Z","shell.execute_reply":"2026-04-16T19:13:02.283283Z"}},"outputs":[],"execution_count":null}]}