{"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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10200,"databundleVersionId":868375,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# CycleGAN: Unpaired Sketch ↔ Photo Translation\n**GenAI Assignment 03 — Question 3**\n\nDomain adaptation using CycleGAN with ResNet generators + PatchGAN discriminators.\n- Domain A: Sketches\n- Domain B: Photos\n- Unpaired training with cycle consistency + identity loss","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Cell 1: Imports & Device Setup\n# ============================================================\nimport os, gc, sys, time, glob, json, random, warnings, itertools, shutil\nimport numpy as np\nfrom PIL import Image, ImageDraw\nfrom pathlib import Path\nfrom collections import deque\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\n\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport torchvision.utils as vutils\n\nimport matplotlib.pyplot as plt\nwarnings.filterwarnings('ignore')\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nnum_gpus = torch.cuda.device_count()\nprint(f'PyTorch {torch.__version__}')\nprint(f'GPUs: {num_gpus}')\nfor i in range(num_gpus):\n    props = torch.cuda.get_device_properties(i)\n    # print(f'  GPU {i}: {props.name} — {props.total_mem / 1e9:.1f} GB')\n    print(f'  GPU {i}: {props.name} ')\n\n# Hyperparameters\nIMG_SIZE        = 128\nBATCH_SIZE      = 4\nNUM_EPOCHS      = 10\nDECAY_START     = 5        # LR decay begins here\nLR              = 2e-4\nBETAS           = (0.5, 0.999)\nLAMBDA_CYCLE    = 10\nLAMBDA_IDENTITY = 5\nBUFFER_SIZE     = 50\nMAX_IMAGES      = 5000      # per domain\nNUM_WORKERS     = 2\nCKPT_EVERY      = 1\nN_RESBLOCKS     = 6\n\nSAVE_DIR = '/kaggle/working'\nos.makedirs(SAVE_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:47.454961Z","iopub.execute_input":"2026-04-08T10:48:47.455519Z","iopub.status.idle":"2026-04-08T10:48:56.328213Z","shell.execute_reply.started":"2026-04-08T10:48:47.455486Z","shell.execute_reply":"2026-04-08T10:48:56.327429Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 2: Explore Available Datasets\n# ============================================================\ndef print_tree(path, prefix='', max_depth=3, _depth=0):\n    if _depth >= max_depth:\n        return\n    try:\n        entries = sorted(os.listdir(path))[:20]  # limit\n        for i, entry in enumerate(entries):\n            full = os.path.join(path, entry)\n            connector = '└── ' if i == len(entries) - 1 else '├── '\n            if os.path.isdir(full):\n                count = len(os.listdir(full))\n                print(f'{prefix}{connector}{entry}/ ({count} items)')\n                ext = '    ' if i == len(entries) - 1 else '│   '\n                print_tree(full, prefix + ext, max_depth, _depth + 1)\n            else:\n                sz = os.path.getsize(full) / 1e6\n                print(f'{prefix}{connector}{entry} ({sz:.1f} MB)')\n    except PermissionError:\n        pass\n\nprint('=== /kaggle/input/ ===')\nfor d in sorted(os.listdir('/kaggle/input')):\n    print(f'\\n--- {d} ---')\n    print_tree(os.path.join('/kaggle/input', d), max_depth=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:56.329695Z","iopub.execute_input":"2026-04-08T10:48:56.330005Z","iopub.status.idle":"2026-04-08T10:48:56.474048Z","shell.execute_reply.started":"2026-04-08T10:48:56.329982Z","shell.execute_reply":"2026-04-08T10:48:56.473316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 3: Locate Sketch and Photo Directories\n# ============================================================\n# Strategy:\n# 1. Try Sketchy Dataset (has both sketches and photos in separate folders)\n# 2. Fallback: TU-Berlin sketches + Anime Faces as photos\n# 3. Fallback: QuickDraw sketches rendered from CSV\n\nSKETCH_DIR = None\nPHOTO_DIR = None\nSKETCHY_BASE = '/kaggle/input/sketchy-dataset'\n\n# --- Attempt 1: Sketchy Dataset ---\n# The Sketchy Dataset typically has: data/photo/ and data/sketch/ or similar\n# Let's find image directories recursively\ndef find_image_dirs(root, min_images=50):\n    \"\"\"Find directories containing images.\"\"\"\n    result = []\n    for dirpath, dirnames, filenames in os.walk(root):\n        imgs = [f for f in filenames if f.lower().endswith(('.png', '.jpg', '.jpeg'))]\n        if len(imgs) >= min_images:\n            result.append((dirpath, len(imgs)))\n    return result\n\nif os.path.exists(SKETCHY_BASE):\n    img_dirs = find_image_dirs(SKETCHY_BASE, min_images=10)\n    print('Image directories in Sketchy Dataset:')\n    for d, c in sorted(img_dirs, key=lambda x: -x[1])[:20]:\n        print(f'  {d} — {c} images')\nelse:\n    print('Sketchy dataset not found at', SKETCHY_BASE)\n\nprint('\\n--- You may need to set SKETCH_DIR and PHOTO_DIR manually in the next cell ---')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:56.475034Z","iopub.execute_input":"2026-04-08T10:48:56.475382Z","iopub.status.idle":"2026-04-08T10:48:56.483401Z","shell.execute_reply.started":"2026-04-08T10:48:56.475357Z","shell.execute_reply":"2026-04-08T10:48:56.482483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 4: Configure Data Sources\n# ============================================================\n# >>> EDIT THESE PATHS based on Cell 2 & 3 output <<<\n#\n# OPTION A — Sketchy Dataset (if it has sketch + photo folders):\n# SKETCH_DIR = '/kaggle/input/sketchy-dataset/data/sketch'\n# PHOTO_DIR  = '/kaggle/input/sketchy-dataset/data/photo'\n#\n# OPTION B — CUHK Sketches + Anime Faces:\n# SKETCH_DIR = '/kaggle/input/cuhk-face-sketch-database-cufs/cropped_sketch'\n# PHOTO_DIR  = '/kaggle/input/anime-faces/data'\n#\n# OPTION C — TU-Berlin (loaded from HF) + any photo dataset:\n# We'll handle TU-Berlin with HuggingFace loader below.\n#\n# OPTION D — QuickDraw rendered sketches + photos from Sketchy:\n# We render QuickDraw CSVs into images.\n\nUSE_QUICKDRAW_SKETCHES = False  # Set True if no sketch folder found\nUSE_HUGGINGFACE_TUBERLIN = False  # Set True to load TU-Berlin from HF\n\n# Default: try Sketchy Dataset paths — ADJUST AFTER INSPECTING Cell 2/3\nSKETCH_DIR = None  # will be set below\nPHOTO_DIR = None\n\n# Auto-detect from Sketchy dataset\nsketchy_base = '/kaggle/input/sketchy-dataset'\nif os.path.exists(sketchy_base):\n    # Look for directories with 'sketch' or 'photo' in name\n    for dirpath, dirnames, filenames in os.walk(sketchy_base):\n        dirname_lower = os.path.basename(dirpath).lower()\n        imgs = [f for f in filenames if f.lower().endswith(('.png', '.jpg', '.jpeg'))]\n        if len(imgs) > 50:\n            if 'sketch' in dirname_lower and SKETCH_DIR is None:\n                SKETCH_DIR = dirpath\n            elif 'photo' in dirname_lower and PHOTO_DIR is None:\n                PHOTO_DIR = dirpath\n\n# Fallback: CUHK sketches + Anime faces\ncuhk_path = '/kaggle/input/cuhk-face-sketch-database-cufs/cropped_sketch'\nanime_path = '/kaggle/input/anime-faces/data'\n\nif SKETCH_DIR is None and os.path.exists(cuhk_path):\n    SKETCH_DIR = cuhk_path\n    print(f'Using CUHK sketches: {cuhk_path}')\nif PHOTO_DIR is None and os.path.exists(anime_path):\n    PHOTO_DIR = anime_path\n    print(f'Using Anime Faces photos: {anime_path}')\n\n# QuickDraw fallback for sketches\nquickdraw_path = '/kaggle/input/competitions/quickdraw-doodle-recognition/train_simplified'\nif SKETCH_DIR is None and os.path.exists(quickdraw_path):\n    USE_QUICKDRAW_SKETCHES = True\n    print('Will render QuickDraw CSVs as sketches')\n\n# TU-Berlin HuggingFace fallback\nif SKETCH_DIR is None and not USE_QUICKDRAW_SKETCHES:\n    USE_HUGGINGFACE_TUBERLIN = True\n    print('Will load TU-Berlin from HuggingFace')\n\nprint(f'\\nFinal config:')\nprint(f'  SKETCH_DIR: {SKETCH_DIR}')\nprint(f'  PHOTO_DIR:  {PHOTO_DIR}')\nprint(f'  USE_QUICKDRAW: {USE_QUICKDRAW_SKETCHES}')\nprint(f'  USE_HF_TUBERLIN: {USE_HUGGINGFACE_TUBERLIN}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:56.485291Z","iopub.execute_input":"2026-04-08T10:48:56.485595Z","iopub.status.idle":"2026-04-08T10:48:56.50276Z","shell.execute_reply.started":"2026-04-08T10:48:56.485573Z","shell.execute_reply":"2026-04-08T10:48:56.501941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 5: Dataset Classes\n# ============================================================\n\ndef collect_image_paths(root, max_count=None, recursive=True):\n    \"\"\"Collect image file paths from a directory.\"\"\"\n    exts = {'.png', '.jpg', '.jpeg', '.webp', '.bmp'}\n    paths = []\n    if recursive:\n        for dirpath, _, filenames in os.walk(root):\n            for f in filenames:\n                if os.path.splitext(f)[1].lower() in exts:\n                    paths.append(os.path.join(dirpath, f))\n    else:\n        for f in os.listdir(root):\n            if os.path.splitext(f)[1].lower() in exts:\n                paths.append(os.path.join(root, f))\n    paths.sort()\n    if max_count and len(paths) > max_count:\n        random.shuffle(paths)\n        paths = paths[:max_count]\n    return paths\n\n\nclass ImageFolderDataset(Dataset):\n    \"\"\"Simple dataset that loads images from a list of paths.\"\"\"\n    def __init__(self, image_paths, img_size=128, is_train=True, convert_mode='RGB'):\n        self.paths = image_paths\n        self.img_size = img_size\n        self.is_train = is_train\n        self.convert_mode = convert_mode\n    \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, idx):\n        img = Image.open(self.paths[idx]).convert(self.convert_mode)\n        img = TF.resize(img, [self.img_size, self.img_size],\n                         interpolation=T.InterpolationMode.BICUBIC)\n        if self.is_train and random.random() > 0.5:\n            img = TF.hflip(img)\n        img = TF.to_tensor(img) * 2.0 - 1.0  # [-1, 1]\n        return img\n\n\nclass QuickDrawDataset(Dataset):\n    \"\"\"Renders Google QuickDraw strokes from CSV into images.\"\"\"\n    def __init__(self, csv_dir, categories=None, per_cat=2000, img_size=128, is_train=True):\n        import pandas as pd\n        if categories is None:\n            categories = ['cat', 'dog', 'house', 'car', 'tree']\n        self.img_size = img_size\n        self.is_train = is_train\n        self.drawings = []\n        \n        for cat in categories:\n            csv_file = os.path.join(csv_dir, f'{cat}.csv')\n            if not os.path.exists(csv_file):\n                print(f'  QuickDraw CSV not found: {csv_file}, skipping')\n                continue\n            df = pd.read_csv(csv_file, usecols=['drawing'], nrows=per_cat)\n            for _, row in df.iterrows():\n                self.drawings.append(json.loads(row['drawing']))\n        \n        random.shuffle(self.drawings)\n        print(f'  QuickDraw: loaded {len(self.drawings)} drawings')\n    \n    def __len__(self):\n        return len(self.drawings)\n    \n    def _render(self, strokes):\n        img = Image.new('RGB', (256, 256), 'white')\n        draw = ImageDraw.Draw(img)\n        for stroke in strokes:\n            xs, ys = stroke[0], stroke[1]\n            points = list(zip(xs, ys))\n            if len(points) > 1:\n                draw.line(points, fill='black', width=3)\n        return img.resize((self.img_size, self.img_size), Image.BICUBIC)\n    \n    def __getitem__(self, idx):\n        img = self._render(self.drawings[idx])\n        if self.is_train and random.random() > 0.5:\n            img = TF.hflip(img)\n        img = TF.to_tensor(img) * 2.0 - 1.0\n        return img\n\n\nclass HuggingFaceSketchDataset(Dataset):\n    \"\"\"Loads TU-Berlin sketches from HuggingFace datasets.\"\"\"\n    def __init__(self, hf_dataset, img_size=128, max_count=5000, is_train=True):\n        self.dataset = hf_dataset\n        self.img_size = img_size\n        self.is_train = is_train\n        self.indices = list(range(min(len(hf_dataset), max_count)))\n        random.shuffle(self.indices)\n    \n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n        item = self.dataset[self.indices[idx]]\n        img = item['image'].convert('RGB')\n        img = TF.resize(img, [self.img_size, self.img_size],\n                         interpolation=T.InterpolationMode.BICUBIC)\n        if self.is_train and random.random() > 0.5:\n            img = TF.hflip(img)\n        img = TF.to_tensor(img) * 2.0 - 1.0\n        return img\n\n\nprint('Dataset classes defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:56.503882Z","iopub.execute_input":"2026-04-08T10:48:56.504206Z","iopub.status.idle":"2026-04-08T10:48:56.52274Z","shell.execute_reply.started":"2026-04-08T10:48:56.504184Z","shell.execute_reply":"2026-04-08T10:48:56.521946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n#  CELL 5.5 — Download & Extract Sketchy Photos (Domain B)\n# ============================================================\nimport os\n\nPHOTO_OUTPUT_DIR = '/kaggle/working/sketchy_photos'\n\nif not os.path.exists(PHOTO_OUTPUT_DIR):\n    print(\"🚀 Downloading 4GB Sketchy Photos from Google Drive...\")\n    \n    # 1. Install gdown (quietly)\n    !pip install gdown -q\n    \n    # 2. Download the file using the specific ID and resourcekey you provided\n    !gdown \"https://drive.google.com/uc?id=0B7ISyeE8QtDdTjE1MG9Gcy1kSkE&export=download&resourcekey=0-r6nB4crmdU-LK7H38xnOUw\" -O sketchy_photos.7z\n    \n    print(\"\\n📦 Extracting .7z file (this might take 1-2 minutes)...\")\n    # 3. Create directory and extract\n    !mkdir -p {PHOTO_OUTPUT_DIR}\n    !7z x sketchy_photos.7z -o{PHOTO_OUTPUT_DIR} -y > /dev/null\n    \n    print(\"🧹 Cleaning up compressed file to save disk space...\")\n    # 4. Delete the heavy .7z file now that we have the images\n    !rm sketchy_photos.7z\n    print(\"✅ Done! Photos are ready.\")\nelse:\n    print(\"✅ Photos already downloaded and extracted in /kaggle/working/sketchy_photos\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:48:56.523561Z","iopub.execute_input":"2026-04-08T10:48:56.523843Z","iopub.status.idle":"2026-04-08T10:55:05.578843Z","shell.execute_reply.started":"2026-04-08T10:48:56.523815Z","shell.execute_reply":"2026-04-08T10:55:05.577799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Domain B: Photos ---\nPHOTO_DIR = '/kaggle/working/sketchy_photos/256x256/photo'\nSKETCH_DIR = '/kaggle/working/sketchy_photos/256x256/sketch'\nprint(f'Loading photos from: {PHOTO_DIR}')\nphoto_paths = collect_image_paths(PHOTO_DIR, MAX_IMAGES, recursive=True)\ndataset_B = ImageFolderDataset(photo_paths, IMG_SIZE)\n\n# ============================================================\n# Cell 6: Build DataLoaders\n# ============================================================\n\n# --- Domain A: Sketches ---\nif USE_HUGGINGFACE_TUBERLIN:\n    print('Loading TU-Berlin from HuggingFace...')\n    from datasets import load_dataset\n    hf_ds = load_dataset('sdiaeyu6n/tu-berlin', split='train')\n    dataset_A = HuggingFaceSketchDataset(hf_ds, IMG_SIZE, MAX_IMAGES)\nelif USE_QUICKDRAW_SKETCHES:\n    print('Rendering QuickDraw sketches...')\n    dataset_A = QuickDrawDataset(quickdraw_path, per_cat=1000, img_size=IMG_SIZE)\nelse:\n    print(f'Loading sketches from: {SKETCH_DIR}')\n    sketch_paths = collect_image_paths(SKETCH_DIR, MAX_IMAGES, recursive=True)\n    dataset_A = ImageFolderDataset(sketch_paths, IMG_SIZE)\n\n# --- Domain B: Photos ---\nprint(f'Loading photos from: {PHOTO_DIR}')\nphoto_paths = collect_image_paths(PHOTO_DIR, MAX_IMAGES, recursive=True)\ndataset_B = ImageFolderDataset(photo_paths, IMG_SIZE)\n\nloader_A = DataLoader(dataset_A, batch_size=BATCH_SIZE, shuffle=True,\n                      num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\nloader_B = DataLoader(dataset_B, batch_size=BATCH_SIZE, shuffle=True,\n                      num_workers=NUM_WORKERS, pin_memory=True, drop_last=True)\n\nprint(f'Domain A (Sketches): {len(dataset_A)} images, {len(loader_A)} batches')\nprint(f'Domain B (Photos):   {len(dataset_B)} images, {len(loader_B)} batches')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:05.580707Z","iopub.execute_input":"2026-04-08T10:55:05.581037Z","iopub.status.idle":"2026-04-08T10:55:06.896355Z","shell.execute_reply.started":"2026-04-08T10:55:05.580985Z","shell.execute_reply":"2026-04-08T10:55:06.895526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ============================================================\n# # Cell 7: Visualize Samples from Both Domains\n# # ============================================================\n# def denorm(x):\n#     return (x + 1.0) / 2.0\n\n# fig, axes = plt.subplots(2, 8, figsize=(20, 5))\n# samples_A = next(iter(loader_A))[:8]\n# samples_B = next(iter(loader_B))[:8]\n\n# for i in range(8):\n#     axes[0, i].imshow(denorm(samples_A[i]).permute(1, 2, 0).numpy().clip(0, 1))\n#     axes[0, i].axis('off')\n#     axes[1, i].imshow(denorm(samples_B[i]).permute(1, 2, 0).numpy().clip(0, 1))\n#     axes[1, i].axis('off')\n\n# axes[0, 0].set_ylabel('Domain A\\n(Sketch)', fontsize=12)\n# axes[1, 0].set_ylabel('Domain B\\n(Photo)', fontsize=12)\n# plt.suptitle('Training Samples', fontsize=14)\n# plt.tight_layout()\n# plt.savefig(os.path.join(SAVE_DIR, 'domain_samples.png'), dpi=100)\n# plt.show()\n# ============================================================\n# Cell 7: Visualize Samples from Both Domains\n# ============================================================\nimport os\nimport matplotlib.pyplot as plt\n\ndef denorm(x):\n    return (x + 1.0) / 2.0\n\n# Fetch one batch of data\nsamples_A = next(iter(loader_A))\nsamples_B = next(iter(loader_B))\n\n# Dynamically set the number of samples to plot (max 8, or whatever the batch size is)\nnum_samples = min(8, len(samples_A), len(samples_B))\n\n# Create subplots based on num_samples\nfig, axes = plt.subplots(2, num_samples, figsize=(2.5 * num_samples, 5))\n\nfor i in range(num_samples):\n    # Handle the array indexing safely\n    ax_A = axes[0, i] if num_samples > 1 else axes[0]\n    ax_B = axes[1, i] if num_samples > 1 else axes[1]\n    \n    # .cpu() added for safety in case tensors are on the GPU\n    ax_A.imshow(denorm(samples_A[i]).cpu().permute(1, 2, 0).numpy().clip(0, 1))\n    ax_A.axis('off')\n    \n    ax_B.imshow(denorm(samples_B[i]).cpu().permute(1, 2, 0).numpy().clip(0, 1))\n    ax_B.axis('off')\n\n# Set side labels for the first column\nif num_samples > 1:\n    axes[0, 0].set_ylabel('Domain A\\n(Sketch)', fontsize=12, labelpad=20, rotation=0, ha='right', va='center')\n    axes[1, 0].set_ylabel('Domain B\\n(Photo)', fontsize=12, labelpad=20, rotation=0, ha='right', va='center')\n\nplt.suptitle(f'Training Samples (Batch Size: {num_samples})', fontsize=14)\nplt.tight_layout()\nOUT_DIR = '/kaggle/working/'\n# Make sure SAVE_DIR exists, otherwise use OUT_DIR\nsave_path = os.path.join(OUT_DIR, 'domain_samples.png')\nplt.savefig(save_path, dpi=100)\nprint(f\"Saved visualization to {save_path}\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:06.8974Z","iopub.execute_input":"2026-04-08T10:55:06.897719Z","iopub.status.idle":"2026-04-08T10:55:08.00445Z","shell.execute_reply.started":"2026-04-08T10:55:06.897688Z","shell.execute_reply":"2026-04-08T10:55:08.00361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 8: ResNet Generator\n# ============================================================\n\nclass ResNetBlock(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.block = nn.Sequential(\n            nn.ReflectionPad2d(1),\n            nn.Conv2d(channels, channels, 3, bias=False),\n            nn.InstanceNorm2d(channels),\n            nn.ReLU(inplace=True),\n            nn.ReflectionPad2d(1),\n            nn.Conv2d(channels, channels, 3, bias=False),\n            nn.InstanceNorm2d(channels),\n        )\n    \n    def forward(self, x):\n        return x + self.block(x)\n\n\nclass ResNetGenerator(nn.Module):\n    \"\"\"ResNet-based generator with reflection padding and InstanceNorm.\"\"\"\n    def __init__(self, in_ch=3, out_ch=3, ngf=64, n_blocks=6):\n        super().__init__()\n        # Initial convolution\n        model = [\n            nn.ReflectionPad2d(3),\n            nn.Conv2d(in_ch, ngf, 7, bias=False),\n            nn.InstanceNorm2d(ngf),\n            nn.ReLU(inplace=True),\n        ]\n        # Downsampling\n        in_features = ngf\n        for _ in range(2):\n            out_features = in_features * 2\n            model += [\n                nn.Conv2d(in_features, out_features, 3, stride=2, padding=1, bias=False),\n                nn.InstanceNorm2d(out_features),\n                nn.ReLU(inplace=True),\n            ]\n            in_features = out_features\n        # ResNet blocks\n        for _ in range(n_blocks):\n            model.append(ResNetBlock(in_features))\n        # Upsampling\n        for _ in range(2):\n            out_features = in_features // 2\n            model += [\n                nn.ConvTranspose2d(in_features, out_features, 3, stride=2,\n                                   padding=1, output_padding=1, bias=False),\n                nn.InstanceNorm2d(out_features),\n                nn.ReLU(inplace=True),\n            ]\n            in_features = out_features\n        # Output\n        model += [\n            nn.ReflectionPad2d(3),\n            nn.Conv2d(ngf, out_ch, 7),\n            nn.Tanh(),\n        ]\n        self.model = nn.Sequential(*model)\n    \n    def forward(self, x):\n        return self.model(x)\n\n\n# Shape test\n_g = ResNetGenerator(n_blocks=N_RESBLOCKS)\n_o = _g(torch.randn(1, 3, 128, 128))\nprint(f'Generator output: {_o.shape}')  # [1, 3, 128, 128]\nprint(f'Generator params: {sum(p.numel() for p in _g.parameters()):,}')\ndel _g, _o","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:08.005974Z","iopub.execute_input":"2026-04-08T10:55:08.006207Z","iopub.status.idle":"2026-04-08T10:55:08.519806Z","shell.execute_reply.started":"2026-04-08T10:55:08.006182Z","shell.execute_reply":"2026-04-08T10:55:08.51905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 9: PatchGAN Discriminator\n# ============================================================\n\nclass PatchGANDiscriminator(nn.Module):\n    \"\"\"PatchGAN with InstanceNorm. No sigmoid — using MSE loss (LSGAN).\"\"\"\n    def __init__(self, in_ch=3, ndf=64):\n        super().__init__()\n        self.model = nn.Sequential(\n            # C64 — no norm\n            nn.Conv2d(in_ch, ndf, 4, stride=2, padding=1),\n            nn.LeakyReLU(0.2, inplace=True),\n            # C128\n            nn.Conv2d(ndf, ndf * 2, 4, stride=2, padding=1, bias=False),\n            nn.InstanceNorm2d(ndf * 2),\n            nn.LeakyReLU(0.2, inplace=True),\n            # C256\n            nn.Conv2d(ndf * 2, ndf * 4, 4, stride=2, padding=1, bias=False),\n            nn.InstanceNorm2d(ndf * 4),\n            nn.LeakyReLU(0.2, inplace=True),\n            # C512\n            nn.Conv2d(ndf * 4, ndf * 8, 4, stride=1, padding=1, bias=False),\n            nn.InstanceNorm2d(ndf * 8),\n            nn.LeakyReLU(0.2, inplace=True),\n            # Output 1 channel\n            nn.Conv2d(ndf * 8, 1, 4, stride=1, padding=1),\n        )\n    \n    def forward(self, x):\n        return self.model(x)\n\n\n_d = PatchGANDiscriminator()\n_o = _d(torch.randn(1, 3, 128, 128))\nprint(f'Discriminator output: {_o.shape}')  # [1, 1, 14, 14] or similar\nprint(f'Discriminator params: {sum(p.numel() for p in _d.parameters()):,}')\ndel _d, _o","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:08.521987Z","iopub.execute_input":"2026-04-08T10:55:08.5222Z","iopub.status.idle":"2026-04-08T10:55:08.579864Z","shell.execute_reply.started":"2026-04-08T10:55:08.522181Z","shell.execute_reply":"2026-04-08T10:55:08.579192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 10: Replay Buffer\n# ============================================================\n\nclass ReplayBuffer:\n    \"\"\"Stores previously generated images for discriminator training.\"\"\"\n    def __init__(self, max_size=50):\n        self.max_size = max_size\n        self.data = []\n    \n    def push_and_pop(self, images):\n        \"\"\"Add new images, return mix of new and buffered.\"\"\"\n        result = []\n        for img in images:\n            img = img.unsqueeze(0)  # [1, C, H, W]\n            if len(self.data) < self.max_size:\n                self.data.append(img.clone())\n                result.append(img)\n            else:\n                if random.random() > 0.5:\n                    idx = random.randint(0, self.max_size - 1)\n                    old = self.data[idx].clone()\n                    self.data[idx] = img.clone()\n                    result.append(old)\n                else:\n                    result.append(img)\n        return torch.cat(result, dim=0)\n\n\nprint('ReplayBuffer ready.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:08.580629Z","iopub.execute_input":"2026-04-08T10:55:08.580928Z","iopub.status.idle":"2026-04-08T10:55:08.587275Z","shell.execute_reply.started":"2026-04-08T10:55:08.580896Z","shell.execute_reply":"2026-04-08T10:55:08.586627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 11: Initialize Models, Losses, Optimizers, Schedulers\n# ============================================================\n\ndef init_weights(m):\n    classname = m.__class__.__name__\n    if classname.find('Conv') != -1:\n        nn.init.normal_(m.weight.data, 0.0, 0.02)\n        if hasattr(m, 'bias') and m.bias is not None:\n            nn.init.constant_(m.bias.data, 0.0)\n    elif classname.find('InstanceNorm') != -1:\n        if m.weight is not None:\n            nn.init.normal_(m.weight.data, 1.0, 0.02)\n            nn.init.constant_(m.bias.data, 0.0)\n\n# Create networks\nG_AB = ResNetGenerator(n_blocks=N_RESBLOCKS)  # Sketch -> Photo\nG_BA = ResNetGenerator(n_blocks=N_RESBLOCKS)  # Photo -> Sketch\nD_A = PatchGANDiscriminator()                  # Is it a real sketch?\nD_B = PatchGANDiscriminator()                  # Is it a real photo?\n\nG_AB.apply(init_weights)\nG_BA.apply(init_weights)\nD_A.apply(init_weights)\nD_B.apply(init_weights)\n\n# DataParallel\nif num_gpus > 1:\n    G_AB = nn.DataParallel(G_AB)\n    G_BA = nn.DataParallel(G_BA)\n    D_A = nn.DataParallel(D_A)\n    D_B = nn.DataParallel(D_B)\n    print(f'DataParallel on {num_gpus} GPUs')\n\nG_AB = G_AB.to(device)\nG_BA = G_BA.to(device)\nD_A = D_A.to(device)\nD_B = D_B.to(device)\n\n# Losses\ncriterion_GAN = nn.MSELoss()   # LSGAN\ncriterion_cycle = nn.L1Loss()\ncriterion_identity = nn.L1Loss()\n\n# Optimizers\noptimizer_G = optim.Adam(\n    itertools.chain(G_AB.parameters(), G_BA.parameters()),\n    lr=LR, betas=BETAS\n)\noptimizer_D = optim.Adam(\n    itertools.chain(D_A.parameters(), D_B.parameters()),\n    lr=LR, betas=BETAS\n)\n\n# LR Schedulers — linear decay to 0 over last half of training\ndef lambda_rule(epoch):\n    return 1.0 - max(0, epoch - DECAY_START) / (NUM_EPOCHS - DECAY_START)\n\nscheduler_G = optim.lr_scheduler.LambdaLR(optimizer_G, lr_lambda=lambda_rule)\nscheduler_D = optim.lr_scheduler.LambdaLR(optimizer_D, lr_lambda=lambda_rule)\n\n# Mixed precision\nscaler_G = GradScaler('cuda')\nscaler_D = GradScaler('cuda')\n\n# Replay buffers\nbuffer_A = ReplayBuffer(BUFFER_SIZE)\nbuffer_B = ReplayBuffer(BUFFER_SIZE)\n\ntotal_params = sum(sum(p.numel() for p in m.parameters()) for m in [G_AB, G_BA, D_A, D_B])\nprint(f'Total parameters: {total_params:,}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:08.588876Z","iopub.execute_input":"2026-04-08T10:55:08.589163Z","iopub.status.idle":"2026-04-08T10:55:08.902827Z","shell.execute_reply.started":"2026-04-08T10:55:08.589141Z","shell.execute_reply":"2026-04-08T10:55:08.902017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 12: Training Loop\n# ============================================================\n\n# Logging\nlog_G, log_D_A, log_D_B, log_cycle = [], [], [], []\n\n# Fixed samples for visualization\nfixed_A = next(iter(loader_A))[:4].to(device)\nfixed_B = next(iter(loader_B))[:4].to(device)\n\nstart_time = time.time()\nn_batches = min(len(loader_A), len(loader_B))\nprint(f'Training: {NUM_EPOCHS} epochs, {n_batches} batches/epoch')\nprint(f'LR decay starts at epoch {DECAY_START}')\nprint('=' * 60)\n\nfor epoch in range(1, NUM_EPOCHS + 1):\n    G_AB.train(); G_BA.train(); D_A.train(); D_B.train()\n    \n    ep_G, ep_DA, ep_DB, ep_cyc = 0.0, 0.0, 0.0, 0.0\n    \n    for batch_idx, (real_A, real_B) in enumerate(zip(loader_A, loader_B)):\n        real_A = real_A.to(device, non_blocking=True)\n        real_B = real_B.to(device, non_blocking=True)\n        \n        # ==========================================\n        # Train Generators G_AB and G_BA\n        # ==========================================\n        optimizer_G.zero_grad(set_to_none=True)\n        \n        with autocast('cuda'):\n            # Identity loss\n            id_A = G_BA(real_A)  # G_BA should be identity for domain A\n            loss_id_A = criterion_identity(id_A, real_A) * LAMBDA_IDENTITY\n            id_B = G_AB(real_B)  # G_AB should be identity for domain B\n            loss_id_B = criterion_identity(id_B, real_B) * LAMBDA_IDENTITY\n            \n            # GAN loss\n            fake_B = G_AB(real_A)          # Sketch -> Photo\n            pred_fake_B = D_B(fake_B)\n            loss_GAN_AB = criterion_GAN(pred_fake_B, torch.ones_like(pred_fake_B))\n            \n            fake_A = G_BA(real_B)          # Photo -> Sketch\n            pred_fake_A = D_A(fake_A)\n            loss_GAN_BA = criterion_GAN(pred_fake_A, torch.ones_like(pred_fake_A))\n            \n            # Cycle consistency loss\n            recov_A = G_BA(fake_B)         # Sketch -> Photo -> Sketch\n            loss_cycle_A = criterion_cycle(recov_A, real_A) * LAMBDA_CYCLE\n            recov_B = G_AB(fake_A)         # Photo -> Sketch -> Photo\n            loss_cycle_B = criterion_cycle(recov_B, real_B) * LAMBDA_CYCLE\n            \n            # Total generator loss\n            loss_G = (loss_GAN_AB + loss_GAN_BA +\n                      loss_cycle_A + loss_cycle_B +\n                      loss_id_A + loss_id_B)\n        \n        scaler_G.scale(loss_G).backward()\n        scaler_G.step(optimizer_G)\n        scaler_G.update()\n        \n        # ==========================================\n        # Train Discriminator D_A (sketch domain)\n        # ==========================================\n        optimizer_D.zero_grad(set_to_none=True)\n        \n        # fake_A_buf = buffer_A.push_and_pop(fake_A.detach().cpu()).to(device)\n        fake_A_buf = buffer_A.push_and_pop(fake_A.detach())\n        \n        with autocast('cuda'):\n            pred_real_A = D_A(real_A)\n            loss_DA_real = criterion_GAN(pred_real_A, torch.ones_like(pred_real_A))\n            pred_fake_A2 = D_A(fake_A_buf)\n            loss_DA_fake = criterion_GAN(pred_fake_A2, torch.zeros_like(pred_fake_A2))\n            loss_DA = (loss_DA_real + loss_DA_fake) * 0.5\n        \n        # Train Discriminator D_B (photo domain)\n        # fake_B_buf = buffer_B.push_and_pop(fake_B.detach().cpu()).to(device)\n        fake_B_buf = buffer_B.push_and_pop(fake_B.detach())\n        \n        with autocast('cuda'):\n            pred_real_B = D_B(real_B)\n            loss_DB_real = criterion_GAN(pred_real_B, torch.ones_like(pred_real_B))\n            pred_fake_B2 = D_B(fake_B_buf)\n            loss_DB_fake = criterion_GAN(pred_fake_B2, torch.zeros_like(pred_fake_B2))\n            loss_DB = (loss_DB_real + loss_DB_fake) * 0.5\n        \n        loss_D_total = loss_DA + loss_DB\n        scaler_D.scale(loss_D_total).backward()\n        scaler_D.step(optimizer_D)\n        scaler_D.update()\n        \n        # Accumulate\n        ep_G += loss_G.item()\n        ep_DA += loss_DA.item()\n        ep_DB += loss_DB.item()\n        ep_cyc += (loss_cycle_A + loss_cycle_B).item()\n        \n        # Memory cleanup every 10 batches\n        # if batch_idx % 10 == 0:\n            # torch.cuda.empty_cache()\n    \n    # Step schedulers\n    scheduler_G.step()\n    scheduler_D.step()\n    \n    # Log\n    log_G.append(ep_G / n_batches)\n    log_D_A.append(ep_DA / n_batches)\n    log_D_B.append(ep_DB / n_batches)\n    log_cycle.append(ep_cyc / n_batches)\n    \n    elapsed = time.time() - start_time\n    lr_now = scheduler_G.get_last_lr()[0]\n    print(f'Epoch [{epoch}/{NUM_EPOCHS}] '\n          f'G: {log_G[-1]:.3f}  DA: {log_D_A[-1]:.3f}  DB: {log_D_B[-1]:.3f}  '\n          f'Cyc: {log_cycle[-1]:.3f}  LR: {lr_now:.6f}  Time: {elapsed/60:.1f}m')\n    \n    # Visual samples every CKPT_EVERY epochs or epoch 1\n    if epoch % CKPT_EVERY == 0 or epoch == 1:\n        G_AB.eval(); G_BA.eval()\n        with torch.no_grad():\n            with autocast('cuda'):\n                fB = G_AB(fixed_A)\n                rA = G_BA(fB)\n                fA = G_BA(fixed_B)\n                rB = G_AB(fA)\n        \n        fig, axes = plt.subplots(4, 6, figsize=(18, 12))\n        titles_top = ['Real A', 'Fake B', 'Recon A']\n        titles_bot = ['Real B', 'Fake A', 'Recon B']\n        for i in range(min(2, fixed_A.size(0))):\n            row = i\n            imgs = [fixed_A[i], fB[i], rA[i]]\n            for j, im in enumerate(imgs):\n                axes[row, j].imshow(denorm(im).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n                axes[row, j].set_title(titles_top[j])\n                axes[row, j].axis('off')\n        for i in range(min(2, fixed_B.size(0))):\n            row = i + 2\n            imgs = [fixed_B[i], fA[i], rB[i]]\n            for j, im in enumerate(imgs):\n                axes[row, j].imshow(denorm(im).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n                axes[row, j].set_title(titles_bot[j])\n                axes[row, j].axis('off')\n        # Hide unused cols\n        for r in range(4):\n            for c in range(3, 6):\n                axes[r, c].axis('off')\n        plt.suptitle(f'Epoch {epoch}', fontsize=14)\n        plt.tight_layout()\n        plt.savefig(os.path.join(SAVE_DIR, f'cyclegan_epoch_{epoch:03d}.png'), dpi=80)\n        plt.show()\n        plt.close()\n        G_AB.train(); G_BA.train()\n    \n    # Save checkpoint\n    if epoch % CKPT_EVERY == 0:\n        gab_sd = G_AB.module.state_dict() if num_gpus > 1 else G_AB.state_dict()\n        gba_sd = G_BA.module.state_dict() if num_gpus > 1 else G_BA.state_dict()\n        torch.save(gab_sd, os.path.join(SAVE_DIR, f'G_AB_epoch{epoch}.pth'))\n        torch.save(gba_sd, os.path.join(SAVE_DIR, f'G_BA_epoch{epoch}.pth'))\n        print(f'  >> Checkpoint saved at epoch {epoch}')\n    \n    # torch.cuda.empty_cache()\n\ntotal_time = time.time() - start_time\nprint(f'\\nTraining complete! {total_time/60:.1f} min ({total_time/3600:.2f} hr)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T10:55:08.90391Z","iopub.execute_input":"2026-04-08T10:55:08.904637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 13: Loss Plots\n# ============================================================\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(log_G, label='Generator Total', color='blue')\naxes[0].set_title('Generator Loss')\naxes[0].set_xlabel('Epoch'); axes[0].legend(); axes[0].grid(True, alpha=0.3)\n\naxes[1].plot(log_D_A, label='D_A (Sketch)', color='red')\naxes[1].plot(log_D_B, label='D_B (Photo)', color='orange')\naxes[1].set_title('Discriminator Losses')\naxes[1].set_xlabel('Epoch'); axes[1].legend(); axes[1].grid(True, alpha=0.3)\n\naxes[2].plot(log_cycle, label='Cycle Consistency', color='green')\naxes[2].set_title('Cycle Consistency Loss')\naxes[2].set_xlabel('Epoch'); axes[2].legend(); axes[2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(os.path.join(SAVE_DIR, 'cyclegan_losses.png'), dpi=120)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 14: Evaluation — SSIM & PSNR on Cycle Reconstructions\n# ============================================================\nfrom skimage.metrics import structural_similarity as ssim_fn\nfrom skimage.metrics import peak_signal_noise_ratio as psnr_fn\n\nG_AB.eval(); G_BA.eval()\n\nall_ssim_A, all_psnr_A = [], []\nall_ssim_B, all_psnr_B = [], []\n\neval_batches = 20  # evaluate on first N batches\nwith torch.no_grad():\n    for i, (rA, rB) in enumerate(zip(loader_A, loader_B)):\n        if i >= eval_batches:\n            break\n        rA = rA.to(device)\n        rB = rB.to(device)\n        \n        with autocast('cuda'):\n            recA = G_BA(G_AB(rA))  # cycle A\n            recB = G_AB(G_BA(rB))  # cycle B\n        \n        # Denorm to [0,1]\n        rA_np = denorm(rA).cpu().float().clamp(0,1).permute(0,2,3,1).numpy()\n        recA_np = denorm(recA).cpu().float().clamp(0,1).permute(0,2,3,1).numpy()\n        rB_np = denorm(rB).cpu().float().clamp(0,1).permute(0,2,3,1).numpy()\n        recB_np = denorm(recB).cpu().float().clamp(0,1).permute(0,2,3,1).numpy()\n        \n        for j in range(rA_np.shape[0]):\n            all_ssim_A.append(ssim_fn(rA_np[j], recA_np[j], channel_axis=2, data_range=1.0))\n            all_psnr_A.append(psnr_fn(rA_np[j], recA_np[j], data_range=1.0))\n            all_ssim_B.append(ssim_fn(rB_np[j], recB_np[j], channel_axis=2, data_range=1.0))\n            all_psnr_B.append(psnr_fn(rB_np[j], recB_np[j], data_range=1.0))\n\nprint(f'Cycle A (Sketch->Photo->Sketch): SSIM={np.mean(all_ssim_A):.4f}, PSNR={np.mean(all_psnr_A):.2f} dB')\nprint(f'Cycle B (Photo->Sketch->Photo):  SSIM={np.mean(all_ssim_B):.4f}, PSNR={np.mean(all_psnr_B):.2f} dB')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 15: Visualization — 10 Translation Examples\n# ============================================================\n\nG_AB.eval(); G_BA.eval()\n\n# Collect 5 from each direction\nvis_A = []\nfor batch in loader_A:\n    vis_A.append(batch)\n    if sum(b.size(0) for b in vis_A) >= 5:\n        break\nvis_A = torch.cat(vis_A)[:5].to(device)\n\nvis_B = []\nfor batch in loader_B:\n    vis_B.append(batch)\n    if sum(b.size(0) for b in vis_B) >= 5:\n        break\nvis_B = torch.cat(vis_B)[:5].to(device)\n\nwith torch.no_grad():\n    with autocast('cuda'):\n        fake_B = G_AB(vis_A)\n        recon_A = G_BA(fake_B)\n        fake_A = G_BA(vis_B)\n        recon_B = G_AB(fake_A)\n\nfig, axes = plt.subplots(10, 3, figsize=(12, 35))\n\nfor i in range(5):\n    # A -> B -> A\n    axes[i, 0].imshow(denorm(vis_A[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[i, 0].set_title('Real Sketch (A)')\n    axes[i, 1].imshow(denorm(fake_B[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[i, 1].set_title('Translated Photo (G_AB)')\n    axes[i, 2].imshow(denorm(recon_A[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[i, 2].set_title('Reconstructed Sketch')\n\nfor i in range(5):\n    row = i + 5\n    axes[row, 0].imshow(denorm(vis_B[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[row, 0].set_title('Real Photo (B)')\n    axes[row, 1].imshow(denorm(fake_A[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[row, 1].set_title('Translated Sketch (G_BA)')\n    axes[row, 2].imshow(denorm(recon_B[i]).cpu().float().clamp(0,1).permute(1,2,0).numpy())\n    axes[row, 2].set_title('Reconstructed Photo')\n\nfor ax_row in axes:\n    for ax in ax_row:\n        ax.axis('off')\n\nplt.suptitle('CycleGAN Translation Results', fontsize=14)\nplt.tight_layout()\nplt.savefig(os.path.join(SAVE_DIR, 'cyclegan_results.png'), dpi=100)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Cell 16: Save Final Models\n# ============================================================\n\ngab_sd = G_AB.module.state_dict() if num_gpus > 1 else G_AB.state_dict()\ngba_sd = G_BA.module.state_dict() if num_gpus > 1 else G_BA.state_dict()\n\npath_ab = os.path.join(SAVE_DIR, 'cyclegan_G_AB.pth')\npath_ba = os.path.join(SAVE_DIR, 'cyclegan_G_BA.pth')\ntorch.save(gab_sd, path_ab)\ntorch.save(gba_sd, path_ba)\n\nprint(f'G_AB saved: {path_ab} ({os.path.getsize(path_ab)/1e6:.1f} MB)')\nprint(f'G_BA saved: {path_ba} ({os.path.getsize(path_ba)/1e6:.1f} MB)')\n\n# Save logs\nnp.savez(os.path.join(SAVE_DIR, 'cyclegan_logs.npz'),\n         G=log_G, D_A=log_D_A, D_B=log_D_B, cycle=log_cycle,\n         ssim_A=np.mean(all_ssim_A), psnr_A=np.mean(all_psnr_A),\n         ssim_B=np.mean(all_ssim_B), psnr_B=np.mean(all_psnr_B))\n\nprint(f'\\n--- Summary ---')\nprint(f'Epochs: {NUM_EPOCHS}')\nprint(f'Total time: {total_time/60:.1f} min')\nprint(f'Cycle A SSIM: {np.mean(all_ssim_A):.4f}, PSNR: {np.mean(all_psnr_A):.2f} dB')\nprint(f'Cycle B SSIM: {np.mean(all_ssim_B):.4f}, PSNR: {np.mean(all_psnr_B):.2f} dB')\nprint(f'\\nFiles:')\nfor f in sorted(os.listdir(SAVE_DIR)):\n    sz = os.path.getsize(os.path.join(SAVE_DIR, f)) / 1e6\n    print(f'  {f} ({sz:.1f} MB)')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Gradio App (run locally after downloading .pth files)\n\n```python\nimport gradio as gr\nimport torch\nimport torch.nn as nn\nfrom PIL import Image\nimport torchvision.transforms.functional as TF\nimport numpy as np\n\n# Paste ResNetBlock + ResNetGenerator classes from Cell 8 here\n# ...\n\nG_AB = ResNetGenerator(n_blocks=6)\nG_BA = ResNetGenerator(n_blocks=6)\nG_AB.load_state_dict(torch.load('cyclegan_G_AB.pth', map_location='cpu', weights_only=True))\nG_BA.load_state_dict(torch.load('cyclegan_G_BA.pth', map_location='cpu', weights_only=True))\nG_AB.eval(); G_BA.eval()\n\ndef translate(image, direction):\n    img = Image.fromarray(image).convert('RGB')\n    img = TF.resize(img, [128, 128])\n    t = TF.to_tensor(img) * 2.0 - 1.0\n    t = t.unsqueeze(0)\n    model = G_AB if direction == 'Sketch → Photo' else G_BA\n    with torch.no_grad():\n        out = model(t)\n    out = ((out.squeeze(0).permute(1,2,0) + 1) / 2).clamp(0,1).numpy()\n    return (out * 255).astype('uint8')\n\ngr.Interface(\n    fn=translate,\n    inputs=[gr.Image(label='Input'), gr.Radio(['Sketch → Photo', 'Photo → Sketch'])],\n    outputs=gr.Image(label='Output'),\n    title='CycleGAN: Sketch ↔ Photo',\n).launch(share=True)\n```","metadata":{}}]}