{"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":"tpuV5e8","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14591369,"sourceType":"datasetVersion","datasetId":9276509}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport random\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\nfrom pathlib import Path\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-24T11:12:12.053668Z","iopub.execute_input":"2026-01-24T11:12:12.053823Z","iopub.status.idle":"2026-01-24T11:12:15.249661Z","shell.execute_reply.started":"2026-01-24T11:12:12.053804Z","shell.execute_reply":"2026-01-24T11:12:15.248930Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"All credits for the model to : https://www.tensorflow.org/tutorials/generative/pix2pix\nExcept as otherwise noted, the content of this page is licensed under the Creative Commons Attribution 4.0 License, and code samples are licensed under the Apache 2.0 License. \n\nAuthors : Copyright 2019 The TensorFlow Authors.\n\nCode used here is Licensed under the Apache License, Version 2.0 (the \"License\"); \nhttps://github.com/tensorflow/docs/blob/master/site/en/tutorials/generative/pix2pix.ipynb\nChanges to the code: \n* dataset\n* jitter function\n* resolution changed to 320x320","metadata":{}},{"cell_type":"markdown","source":"a conditional generative adversarial network (cGAN) called pix2pix that learns a mapping from input images to output images","metadata":{}},{"cell_type":"markdown","source":"/kaggle/input/vesuvius-slice-image-label-dataset-1c/single_channel_data/train\n/kaggle/input/vesuvius-slice-image-label-dataset-1c/single_channel_data/test\n/kaggle/input/vesuvius-slice-image-label-dataset-1c/single_channel_data/val","metadata":{}},{"cell_type":"markdown","source":"In Pix2Pix, an \"ignore mask\" is a technique used to focus the network’s learning on specific regions of an image while disregarding others during the loss calculation. This is commonly applied in tasks like inpainting or localized editing, where you only want the model to learn from or modify parts of the image that are missing or corrupted. \nKey Implementation Details\nMasked L1 Loss: Instead of calculating the standard L1 loss (pixel-wise difference) across the entire image, the loss is multiplied by a binary mask. This ensures the generator is only penalized for errors within the \"active\" regions.\nConditioning: Even when using a mask to ignore certain areas, the generator and discriminator typically still receive the full input image to maintain global context and structural consistency.\nEdge Smoothing: Some implementations use an additional \"edge mask loss\" to prevent artifacts where the generated region meets the ignored (original) region. \nCommon Applications\nOverexposure Recovery: Recovering missing details in blown-out areas of a photo by masking those specific regions for reconstruction.\nInpainting: Filling in missing pixels or removing objects; the \"ignore mask\" allows the model to leave the rest of the image untouched.\nMedical Imaging: Segmenting specific abnormalities (like pulmonary issues) while ignoring healthy tissue to improve model precision","metadata":{}},{"cell_type":"code","source":"# =========================\n# 1. IMPORTS\n# =========================\nimport time\nfrom pathlib import Path\nimport numpy as np\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\n\nimport jax\nimport jax.numpy as jnp\nfrom jax import random, value_and_grad, pmap\nfrom flax import linen as nn\nfrom flax.training import train_state\nimport optax\nfrom jax.tree_util import tree_map\n\n\nprint(\"JAX devices:\", jax.devices())\nNUM_DEVICES = len(jax.devices())\n\n# =========================\n# 2. CONFIG & DATA\n# =========================\nIMG_WIDTH = 320\nIMG_HEIGHT = 320\nINPUT_CHANNELS = 2\nTARGET_CHANNELS = 3\nLAMBDA = 25.0\nEPOCHS = 50\nGLOBAL_BATCH_SIZE = 512\nPER_DEVICE_BATCH = GLOBAL_BATCH_SIZE // NUM_DEVICES\n\nPATH_DATA = Path('/kaggle/input/vesuvius-image-slices/single_channel_data/train/')\n\ndef load_and_preprocess(image_path):\n    img = tf.io.read_file(image_path)\n    img = tf.image.decode_png(img, channels=3)\n    img = tf.cast(img, tf.float32)\n\n    w = tf.shape(img)[1] // 2\n    target = img[:, :w, :]\n    source = img[:, w:, :]\n\n    combined = tf.concat([target, source], axis=-1)\n\n    # flips\n    combined = tf.image.random_flip_left_right(combined)\n    combined = tf.image.random_flip_up_down(combined)\n\n    # rotations\n    k = tf.random.uniform([], 0, 4, dtype=tf.int32)\n    combined = tf.image.rot90(combined, k=k)\n\n    # --- SAFE TRANSLATION (no wrap) ---\n    combined = tf.pad(combined, [[20,20],[20,20],[0,0]], mode=\"REFLECT\")\n    combined = tf.image.random_crop(combined, [IMG_HEIGHT, IMG_WIDTH, 6])\n\n    # mild color jitter\n    combined = tf.image.random_brightness(combined, 0.05)\n    combined = tf.image.random_contrast(combined, 0.9, 1.1)\n\n    target = combined[:, :, :3]\n    source = combined[:, :, 3:]\n\n    target = (target / 127.5) - 1.0\n    source = (source / 127.5) - 1.0\n\n    source = source[:, :, 1:3]\n    return source, target\n\ndef get_dataset(batch_size=GLOBAL_BATCH_SIZE):\n    files = sorted([str(p) for p in PATH_DATA.glob('*.png')])\n    print(f\"Found {len(files)} combined images.\")\n    ds = tf.data.Dataset.from_tensor_slices(files)\n    ds = ds.shuffle(len(files))\n    ds = ds.map(load_and_preprocess, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.batch(batch_size, drop_remainder=True)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds\n\ntrain_ds = get_dataset()\n\ndef tf_to_np(a, b):\n    return np.array(a), np.array(b)\n\n# def replicate(x):\n#     return jax.tree_map(lambda v: jnp.stack([v] * NUM_DEVICES), x)\ndef replicate(x):\n    return tree_map(lambda v: jnp.stack([v] * NUM_DEVICES), x)\ndef shard(x):\n    return x.reshape(NUM_DEVICES, -1, *x.shape[1:])\n\n# =========================\n# 3. MODEL DEFINITIONS\n# =========================\nACT = jnp.bfloat16\nPARAM = jnp.float32\n\nclass Down(nn.Module):\n    f: int\n    k: int\n    bn: bool = True\n    @nn.compact\n    def __call__(self, x, train=True):\n        x = nn.Conv(self.f, (self.k,self.k), strides=(2,2),\n                    padding='SAME',\n                    kernel_init=nn.initializers.normal(0.02),\n                    use_bias=not self.bn,\n                    dtype=ACT, param_dtype=PARAM)(x)\n        if self.bn:\n            x = nn.BatchNorm(use_running_average=not train,\n                             dtype=ACT, param_dtype=PARAM)(x)\n        return nn.leaky_relu(x, 0.2)\n\nclass Up(nn.Module):\n    f: int\n    k: int\n    drop: bool = False\n    @nn.compact\n    def __call__(self, x, train=True):\n        x = nn.ConvTranspose(self.f, (self.k,self.k), strides=(2,2),\n                             padding='SAME',\n                             kernel_init=nn.initializers.normal(0.02),\n                             use_bias=False,\n                             dtype=ACT, param_dtype=PARAM)(x)\n        x = nn.BatchNorm(use_running_average=not train,\n                         dtype=ACT, param_dtype=PARAM)(x)\n        if self.drop and train:\n            x = nn.Dropout(0.5)(x, deterministic=not train)\n        return nn.relu(x)\n\nclass Generator(nn.Module):\n    @nn.compact\n    def __call__(self, x, train=True):\n        x = x.astype(ACT)\n        downs = [\n            Down(64,4,False),\n            Down(128,4),\n            Down(256,4),\n            Down(512,4),\n            Down(512,4),\n            Down(512,4),\n        ]\n        ups = [\n            Up(512,4,True),\n            Up(512,4,True),\n            Up(256,4),\n            Up(128,4),\n            Up(64,4),\n        ]\n        skips = []\n        h = x\n        for d in downs:\n            h = d(h, train=train)\n            skips.append(h)\n        for u, s in zip(ups, reversed(skips[:-1])):\n            h = u(h, train=train)\n            if train:\n                s = nn.Dropout(0.1)(s, deterministic=not train)\n            h = jnp.concatenate([h, s], axis=-1)\n        h = nn.ConvTranspose(TARGET_CHANNELS, (4,4), strides=(2,2),\n                             padding='SAME',\n                             kernel_init=nn.initializers.normal(0.02),\n                             dtype=ACT, param_dtype=PARAM)(h)\n        return jnp.tanh(h).astype(jnp.float32)\n\nclass Discriminator(nn.Module):\n    @nn.compact\n    def __call__(self, source, target, train=True):\n        x = jnp.concatenate([source.astype(ACT), target.astype(ACT)], axis=-1)\n        x = Down(64,4,False)(x,train)\n        x = Down(128,4)(x,train)\n        x = Down(256,4)(x,train)\n        x = jnp.pad(x, ((0,0),(1,1),(1,1),(0,0)))\n        x = nn.Conv(512,(4,4),strides=(1,1),padding='VALID',\n                    kernel_init=nn.initializers.normal(0.02),\n                    use_bias=False,\n                    dtype=ACT, param_dtype=PARAM)(x)\n        x = nn.BatchNorm(use_running_average=not train,\n                         dtype=ACT, param_dtype=PARAM)(x)\n        x = nn.leaky_relu(x,0.2)\n        x = jnp.pad(x, ((0,0),(1,1),(1,1),(0,0)))\n        x = nn.Conv(1,(4,4),strides=(1,1),padding='VALID',\n                    kernel_init=nn.initializers.normal(0.02),\n                    dtype=ACT, param_dtype=PARAM)(x)\n        return x.astype(jnp.float32)\n\n# =========================\n# 4. TRAIN STATE & LOSSES\n# =========================\nclass State(train_state.TrainState):\n    batch_stats: dict\n\ndef bce(logits, labels):\n    return optax.sigmoid_binary_cross_entropy(logits, labels).mean()\n\ndef g_loss_fn(g_state, d_state, batch, rng):\n    source, target = batch\n    G = Generator()\n    D = Discriminator()\n\n    vars_g = {'params': g_state.params, 'batch_stats': g_state.batch_stats}\n    (pred, new_g_vars) = G.apply(vars_g, source, train=True,\n                                 rngs={'dropout':rng}, mutable=['batch_stats'])\n    new_g_bs = new_g_vars['batch_stats']\n\n    vars_d = {'params': d_state.params, 'batch_stats': d_state.batch_stats}\n    (disc_fake, _) = D.apply(vars_d, source, pred, train=True, mutable=['batch_stats'])\n\n    gan = bce(disc_fake, jnp.ones_like(disc_fake))\n    l1  = jnp.mean(jnp.abs(target - pred))\n    total = gan + LAMBDA*l1\n    return total, (gan, l1, new_g_bs)\n\ndef d_loss_fn(g_state, d_state, batch):\n    source, target = batch\n    G = Generator()\n    D = Discriminator()\n\n    pred = G.apply({'params':g_state.params,'batch_stats':g_state.batch_stats},\n                   source, train=False)\n\n    vars_d = {'params': d_state.params, 'batch_stats': d_state.batch_stats}\n    (disc_real, new_d_vars_real) = D.apply(vars_d, source, target, train=True, mutable=['batch_stats'])\n    (disc_fake, new_d_vars_fake) = D.apply(\n        {'params':d_state.params,'batch_stats':new_d_vars_real['batch_stats']},\n        source, pred, train=True, mutable=['batch_stats']\n    )\n\n    real = bce(disc_real, jnp.ones_like(disc_real))\n    fake = bce(disc_fake, jnp.zeros_like(disc_fake))\n    return real+fake, new_d_vars_fake['batch_stats']\n\n# =========================\n# 5. INIT\n# =========================\nkey = random.PRNGKey(0)\ndummy_source = jnp.zeros((1,IMG_HEIGHT,IMG_WIDTH,INPUT_CHANNELS))\ndummy_target = jnp.zeros((1,IMG_HEIGHT,IMG_WIDTH,TARGET_CHANNELS))\n\nG = Generator()\nD = Discriminator()\n\ng_vars = G.init(key, dummy_source)\nd_vars = D.init(key, dummy_source, dummy_target)\n\ng_state = State.create(\n    apply_fn=G.apply,\n    params=g_vars['params'],\n    tx=optax.adam(2e-4, b1=0.5),\n    batch_stats=g_vars['batch_stats'],\n)\nd_state = State.create(\n    apply_fn=D.apply,\n    params=d_vars['params'],\n    tx=optax.adam(2e-4, b1=0.5),\n    batch_stats=d_vars['batch_stats'],\n)\n\n# correct replication\ng_state = replicate(g_state)\nd_state = replicate(d_state)\n\n# =========================\n# 6. PMAP TRAIN STEP\n# =========================\ndef step_single(g_state, d_state, batch, rng):\n    source, target = batch\n\n    def g_fn(params, bs):\n        tmp = g_state.replace(params=params, batch_stats=bs)\n        loss, (gan,l1,new_bs) = g_loss_fn(tmp, d_state, (source,target), rng)\n        return loss, (gan,l1,new_bs)\n\n    (g_loss,(gan,l1,new_g_bs)), g_grads = value_and_grad(g_fn,has_aux=True)(\n        g_state.params, g_state.batch_stats\n    )\n    g_state = g_state.apply_gradients(grads=g_grads)\n    g_state = g_state.replace(batch_stats=new_g_bs)\n\n    def d_fn(params, bs):\n        tmp = d_state.replace(params=params, batch_stats=bs)\n        loss, new_bs = d_loss_fn(g_state, tmp, (source,target))\n        return loss, new_bs\n\n    (d_loss,new_d_bs), d_grads = value_and_grad(d_fn,has_aux=True)(\n        d_state.params, d_state.batch_stats\n    )\n    d_state = d_state.apply_gradients(grads=d_grads)\n    d_state = d_state.replace(batch_stats=new_d_bs)\n\n    return g_state, d_state, g_loss, d_loss, gan, l1\n\np_step = pmap(step_single, axis_name='devices')\n\n# =========================\n# 7. WARMUP COMPILATION\n# =========================\ndef warmup_step(g_state, d_state, train_ds):\n    print(\"Running warmup compilation...\")\n    for batch_source, batch_target in train_ds.take(1):\n        src_np, tgt_np = tf_to_np(batch_source, batch_target)\n        src = shard(src_np)\n        tgt = shard(tgt_np)\n        rng = random.PRNGKey(42)\n        subkeys = random.split(rng, NUM_DEVICES)\n        _ = p_step(g_state, d_state, (src, tgt), subkeys)\n        break\n    print(\"Warmup complete.\\n\")\n\nwarmup_step(g_state, d_state, train_ds)\n\n# =========================\n# 8. VISUALIZATION\n# =========================\ndef show_sample(g_host, source, target, title):\n    G = Generator()\n    pred = G.apply(\n        {'params': g_host.params, 'batch_stats': g_host.batch_stats},\n        source, train=False\n    )\n\n    src_vis = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.float32)\n    src_vis[:, :, :2] = source[0]\n\n    tgt_vis = target[0]\n    pred_vis = pred[0]\n\n    plt.figure(figsize=(12,4))\n    for i, img in enumerate([src_vis, tgt_vis, pred_vis]):\n        plt.subplot(1,3,i+1)\n        vis = np.clip(img * 0.5 + 0.5, 0, 1)\n        plt.imshow(vis)\n        plt.title([\"Source (G,R)\", \"Target (GT)\", title][i])\n        plt.axis(\"off\")\n    plt.show()\n\n# =========================\n# 9. TRAIN LOOP\n# =========================\nrng = key\n\nfor epoch in range(EPOCHS):\n    g_losses=[]; d_losses=[]\n    start=time.time()\n\n    for batch_source, batch_target in train_ds:\n        src_np, tgt_np = tf_to_np(batch_source, batch_target)\n        src = shard(src_np)\n        tgt = shard(tgt_np)\n\n        rng, sub = random.split(rng)\n        subkeys = random.split(sub, NUM_DEVICES)\n\n        g_state, d_state, g_loss, d_loss, gan, l1 = p_step(\n            g_state, d_state, (src,tgt), subkeys\n        )\n\n        g_losses.append(np.array(g_loss).mean())\n        d_losses.append(np.array(d_loss).mean())\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}  G={np.mean(g_losses):.4f}  D={np.mean(d_losses):.4f}  Time={time.time()-start:.1f}s\")\n\n    g_host = jax.tree_util.tree_map(lambda x: x[0], g_state)\n    for ex_src, ex_tgt in train_ds.take(1):\n        src_np, tgt_np = tf_to_np(ex_src, ex_tgt)\n        show_sample(g_host, src_np, tgt_np, f\"Epoch {epoch+1}\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-24T11:19:25.429961Z","iopub.execute_input":"2026-01-24T11:19:25.430287Z","execution_failed":"2026-01-24T11:28:49.501Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"alternative https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix","metadata":{}},{"cell_type":"markdown","source":"Augumentation dataset :) from label to something else :)))\n","metadata":{}},{"cell_type":"code","source":"# # =========================\n# # 1. IMPORTS\n# # =========================\n# import os\n# import time\n# from pathlib import Path\n\n# import numpy as np\n# import tensorflow as tf\n# import matplotlib.pyplot as plt\n\n# import jax\n# import jax.numpy as jnp\n# from jax import random, value_and_grad, pmap\n# from flax import linen as nn\n# from flax.training import train_state, checkpoints\n# import optax\n\n# print(\"JAX devices:\", jax.devices())\n# NUM_DEVICES = len(jax.devices())\n\n# # =========================\n# # 2. CONFIG & DATA\n# # =========================\n# IMG_WIDTH = 320\n# IMG_HEIGHT = 320\n# INPUT_CHANNELS = 2\n# OUTPUT_CHANNELS = 3\n# LAMBDA = 100.0\n# EPOCHS = 40\n# GLOBAL_BATCH_SIZE = 32  # must be divisible by NUM_DEVICES\n# PER_DEVICE_BATCH = GLOBAL_BATCH_SIZE // NUM_DEVICES\n\n# CKPT_DIR = \"/kaggle/working/checkpoints_pix2pix_jax\"\n# os.makedirs(CKPT_DIR, exist_ok=True)\n\n# PATH_DATA = Path('/kaggle/input/vesuvius-image-slices/single_channel_data/train/')\n\n# def load_and_preprocess(image_path):\n#     img = tf.io.read_file(image_path)\n#     img = tf.image.decode_png(img, channels=3)\n#     img = tf.cast(img, tf.float32)\n\n#     w = tf.shape(img)[1] // 2\n#     label = img[:, :w, :]\n#     mask = img[:, w:, :]\n\n#     label = tf.image.resize(img_in, [IMG_HEIGHT, IMG_WIDTH])\n#     mask = tf.image.resize(img_tar, [IMG_HEIGHT, IMG_WIDTH])\n\n#     label = (img_in / 127.5) - 1.0\n#      mask = (img_tar / 127.5) - 1.0\n\n#     img_in = img_in[:, :, :2]  # input = left RG\n#     return img_in, img_tar      # target = right RGB\n\n# def get_dataset(batch_size=GLOBAL_BATCH_SIZE):\n#     all_files = sorted([str(p) for p in PATH_DATA.glob('*.png')])\n#     num_samples = len(all_files)\n#     print(f\"Found {num_samples} combined images (640x320).\")\n#     if num_samples == 0:\n#         raise FileNotFoundError(\"No .png files found.\")\n\n#     ds = tf.data.Dataset.from_tensor_slices(all_files)\n#     ds = ds.shuffle(max(num_samples, 1))\n#     ds = ds.map(load_and_preprocess, num_parallel_calls=tf.data.AUTOTUNE)\n#     ds = ds.batch(batch_size, drop_remainder=True)\n#     ds = ds.prefetch(tf.data.AUTOTUNE)\n#     return ds\n\n# train_ds = get_dataset()\n\n# def tf_batch_to_numpy(batch_in, batch_tar):\n#     return np.array(batch_in), np.array(batch_tar)\n\n# def shard_batch(x):\n#     x = x.reshape(NUM_DEVICES, -1, *x.shape[1:])\n#     return x\n\n# # =========================\n# # 3. MODEL DEFINITIONS (bfloat16 activations)\n# # =========================\n# ACT_DTYPE = jnp.bfloat16\n# PARAM_DTYPE = jnp.float32\n\n# class Downsample(nn.Module):\n#     filters: int\n#     size: int\n#     apply_batchnorm: bool = True\n\n#     @nn.compact\n#     def __call__(self, x, train=True):\n#         x = nn.Conv(\n#             self.filters,\n#             (self.size, self.size),\n#             strides=(2, 2),\n#             padding='SAME',\n#             kernel_init=nn.initializers.normal(0.02),\n#             use_bias=not self.apply_batchnorm,\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         if self.apply_batchnorm:\n#             x = nn.BatchNorm(\n#                 use_running_average=not train,\n#                 name='bn',\n#                 dtype=ACT_DTYPE,\n#                 param_dtype=PARAM_DTYPE,\n#             )(x)\n#         x = nn.leaky_relu(x, negative_slope=0.2)\n#         return x\n\n\n# class Upsample(nn.Module):\n#     filters: int\n#     size: int\n#     apply_dropout: bool = False\n\n#     @nn.compact\n#     def __call__(self, x, train=True):\n#         x = nn.ConvTranspose(\n#             self.filters,\n#             (self.size, self.size),\n#             strides=(2, 2),\n#             padding='SAME',\n#             kernel_init=nn.initializers.normal(0.02),\n#             use_bias=False,\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         x = nn.BatchNorm(\n#             use_running_average=not train,\n#             name='bn',\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         if self.apply_dropout and train:\n#             x = nn.Dropout(0.5)(x, deterministic=not train)\n#         x = nn.relu(x)\n#         return x\n\n\n# class GeneratorUNet(nn.Module):\n#     @nn.compact\n#     def __call__(self, x, train=True):\n#         x = x.astype(ACT_DTYPE)\n#         downs = [\n#             Downsample(64, 4, apply_batchnorm=False),\n#             Downsample(128, 4),\n#             Downsample(256, 4),\n#             Downsample(512, 4),\n#             Downsample(512, 4),\n#             Downsample(512, 4),\n#         ]\n\n#         ups = [\n#             Upsample(512, 4, apply_dropout=True),\n#             Upsample(512, 4, apply_dropout=True),\n#             Upsample(256, 4),\n#             Upsample(128, 4),\n#             Upsample(64, 4),\n#         ]\n\n#         skips = []\n#         h = x\n#         for d in downs:\n#             h = d(h, train=train)\n#             skips.append(h)\n\n#         skips_for_ups = list(reversed(skips[:-1]))\n#         for u, s in zip(ups, skips_for_ups):\n#             h = u(h, train=train)\n#             h = jnp.concatenate([h, s], axis=-1)\n\n#         h = nn.ConvTranspose(\n#             OUTPUT_CHANNELS,\n#             (4, 4),\n#             strides=(2, 2),\n#             padding='SAME',\n#             kernel_init=nn.initializers.normal(0.02),\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(h)\n#         h = jnp.tanh(h)\n#         return h.astype(jnp.float32)\n\n\n# def zero_pad(x, pad):\n#     return jnp.pad(x, pad, mode='constant', constant_values=0.0)\n\n\n# class DiscriminatorPatchGAN(nn.Module):\n#     @nn.compact\n#     def __call__(self, inp, tar, train=True):\n#         inp = inp.astype(ACT_DTYPE)\n#         tar = tar.astype(ACT_DTYPE)\n#         x = jnp.concatenate([inp, tar], axis=-1)\n\n#         x = Downsample(64, 4, apply_batchnorm=False)(x, train=train)\n#         x = Downsample(128, 4)(x, train=train)\n#         x = Downsample(256, 4)(x, train=train)\n\n#         x = zero_pad(x, ((0, 0), (1, 1), (1, 1), (0, 0)))\n#         x = nn.Conv(\n#             512,\n#             (4, 4),\n#             strides=(1, 1),\n#             padding='VALID',\n#             kernel_init=nn.initializers.normal(0.02),\n#             use_bias=False,\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         x = nn.BatchNorm(\n#             use_running_average=not train,\n#             name='bn',\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         x = nn.leaky_relu(x, negative_slope=0.2)\n\n#         x = zero_pad(x, ((0, 0), (1, 1), (1, 1), (0, 0)))\n#         x = nn.Conv(\n#             1,\n#             (4, 4),\n#             strides=(1, 1),\n#             padding='VALID',\n#             kernel_init=nn.initializers.normal(0.02),\n#             dtype=ACT_DTYPE,\n#             param_dtype=PARAM_DTYPE,\n#         )(x)\n#         return x.astype(jnp.float32)\n\n# # =========================\n# # 4. TRAIN STATE & LOSSES\n# # =========================\n# class TrainStateWithBN(train_state.TrainState):\n#     batch_stats: dict\n\n# def bce_logits(logits, labels):\n#     return optax.sigmoid_binary_cross_entropy(logits, labels).mean()\n\n# def generator_loss_fn(gen_state, disc_state, batch, rng):\n#     inp, tar = batch\n#     gen = GeneratorUNet()\n#     disc = DiscriminatorPatchGAN()\n\n#     variables = {'params': gen_state.params, 'batch_stats': gen_state.batch_stats}\n#     (gen_out, new_gen_vars) = gen.apply(\n#         variables,\n#         inp,\n#         train=True,\n#         rngs={'dropout': rng},\n#         mutable=['batch_stats'],\n#     )\n#     new_gen_batch_stats = new_gen_vars['batch_stats']\n\n#     disc_vars = {'params': disc_state.params, 'batch_stats': disc_state.batch_stats}\n#     (disc_fake, _) = disc.apply(\n#         disc_vars,\n#         inp,\n#         gen_out,\n#         train=True,\n#         mutable=['batch_stats'],\n#     )\n\n#     gan_loss = bce_logits(disc_fake, jnp.ones_like(disc_fake))\n#     l1_loss = jnp.mean(jnp.abs(tar - gen_out))\n#     total = gan_loss + LAMBDA * l1_loss\n#     return total, (gan_loss, l1_loss, new_gen_batch_stats)\n\n# def discriminator_loss_fn(gen_state, disc_state, batch):\n#     inp, tar = batch\n#     gen = GeneratorUNet()\n#     disc = DiscriminatorPatchGAN()\n\n#     gen_vars = {'params': gen_state.params, 'batch_stats': gen_state.batch_stats}\n#     gen_out = gen.apply(gen_vars, inp, train=False, mutable=False)\n\n#     disc_vars = {'params': disc_state.params, 'batch_stats': disc_state.batch_stats}\n#     (disc_real, new_disc_vars_real) = disc.apply(\n#         disc_vars,\n#         inp,\n#         tar,\n#         train=True,\n#         mutable=['batch_stats'],\n#     )\n#     (disc_fake, new_disc_vars_fake) = disc.apply(\n#         {'params': disc_state.params, 'batch_stats': new_disc_vars_real['batch_stats']},\n#         inp,\n#         gen_out,\n#         train=True,\n#         mutable=['batch_stats'],\n#     )\n\n#     real_loss = bce_logits(disc_real, jnp.ones_like(disc_real))\n#     fake_loss = bce_logits(disc_fake, jnp.zeros_like(disc_fake))\n#     total = real_loss + fake_loss\n#     new_disc_batch_stats = new_disc_vars_fake['batch_stats']\n#     return total, new_disc_batch_stats\n\n# # =========================\n# # 5. INIT + LR SCHEDULE\n# # =========================\n# key = random.PRNGKey(0)\n# dummy_inp = jnp.zeros((1, IMG_HEIGHT, IMG_WIDTH, INPUT_CHANNELS), jnp.float32)\n# dummy_tar = jnp.zeros((1, IMG_HEIGHT, IMG_WIDTH, OUTPUT_CHANNELS), jnp.float32)\n\n# gen = GeneratorUNet()\n# disc = DiscriminatorPatchGAN()\n\n# gen_vars = gen.init(key, dummy_inp)\n# disc_vars = disc.init(key, dummy_inp, dummy_tar)\n\n# gen_params = gen_vars['params']\n# gen_batch_stats = gen_vars.get('batch_stats', {})\n# disc_params = disc_vars['params']\n# disc_batch_stats = disc_vars.get('batch_stats', {})\n\n# steps_per_epoch = sum(1 for _ in train_ds)\n# total_steps = steps_per_epoch * EPOCHS\n\n# lr_schedule = optax.cosine_decay_schedule(\n#     init_value=2e-4,\n#     decay_steps=total_steps,\n#     alpha=0.1,\n# )\n\n# gen_tx = optax.adam(lr_schedule, b1=0.5)\n# disc_tx = optax.adam(lr_schedule, b1=0.5)\n\n# gen_state = TrainStateWithBN.create(\n#     apply_fn=gen.apply,\n#     params=gen_params,\n#     tx=gen_tx,\n#     batch_stats=gen_batch_stats,\n# )\n# disc_state = TrainStateWithBN.create(\n#     apply_fn=disc.apply,\n#     params=disc_params,\n#     tx=disc_tx,\n#     batch_stats=disc_batch_stats,\n# )\n\n# gen_state = jax.device_put_replicated(gen_state, jax.devices())\n# disc_state = jax.device_put_replicated(disc_state, jax.devices())\n\n# print(\"Initialized. Using devices:\", jax.devices())\n\n# # =========================\n# # 6. PMAPPED TRAIN STEP\n# # =========================\n# def train_step_single(gen_state, disc_state, batch, rng):\n#     inp, tar = batch\n#     inp = inp.astype(jnp.float32)\n#     tar = tar.astype(jnp.float32)\n\n#     def g_loss_fn(params, batch_stats):\n#         tmp_state = gen_state.replace(params=params, batch_stats=batch_stats)\n#         loss, (gan_l, l1_l, new_bs) = generator_loss_fn(tmp_state, disc_state, (inp, tar), rng)\n#         return loss, (gan_l, l1_l, new_bs)\n\n#     (g_loss, (gan_l, l1_l, new_gen_bs)), g_grads = value_and_grad(\n#         g_loss_fn, has_aux=True\n#     )(gen_state.params, gen_state.batch_stats)\n\n#     gen_state = gen_state.apply_gradients(grads=g_grads)\n#     gen_state = gen_state.replace(batch_stats=new_gen_bs)\n\n#     def d_loss_fn(params, batch_stats):\n#         tmp_state = disc_state.replace(params=params, batch_stats=batch_stats)\n#         loss, new_bs = discriminator_loss_fn(gen_state, tmp_state, (inp, tar))\n#         return loss, new_bs\n\n#     (d_loss, new_disc_bs), d_grads = value_and_grad(\n#         d_loss_fn, has_aux=True\n#     )(disc_state.params, disc_state.batch_stats)\n\n#     disc_state = disc_state.apply_gradients(grads=d_grads)\n#     disc_state = disc_state.replace(batch_stats=new_disc_bs)\n\n#     return gen_state, disc_state, g_loss, d_loss, gan_l, l1_l\n\n# p_train_step = pmap(\n#     train_step_single,\n#     axis_name='devices',\n#     in_axes=(0, 0, 0, 0),\n#     out_axes=(0, 0, 0, 0, 0, 0),\n# )\n\n# # =========================\n# # 7. CHECKPOINTS & VAL GRID\n# # =========================\n# history = {'gen_loss': [], 'disc_loss': []}\n# rng = key\n\n# def save_ckpt(epoch, gen_state, disc_state):\n#     host_gen = jax.tree_util.tree_map(lambda x: x[0], gen_state)\n#     host_disc = jax.tree_util.tree_map(lambda x: x[0], disc_state)\n#     to_save = {\n#         'gen_state': host_gen,\n#         'disc_state': host_disc,\n#         'epoch': epoch,\n#     }\n#     checkpoints.save_checkpoint(\n#         CKPT_DIR,\n#         to_save,\n#         step=epoch,\n#         overwrite=True,\n#     )\n\n# def load_ckpt_if_exists(gen_state, disc_state):\n#     ckpt = checkpoints.restore_checkpoint(CKPT_DIR, target=None)\n#     if ckpt:\n#         print(\"Loaded checkpoint at epoch\", ckpt['epoch'])\n#         g = ckpt['gen_state']\n#         d = ckpt['disc_state']\n#         gen_state_new = jax.device_put_replicated(gen_state.replace(\n#             params=g['params'],\n#             batch_stats=g['batch_stats'],\n#             opt_state=g['opt_state'],\n#         ), jax.devices())\n#         disc_state_new = jax.device_put_replicated(disc_state.replace(\n#             params=d['params'],\n#             batch_stats=d['batch_stats'],\n#             opt_state=d['opt_state'],\n#         ), jax.devices())\n#         return gen_state_new, disc_state_new, ckpt['epoch']\n#     return gen_state, disc_state, 0\n\n# gen_state, disc_state, start_epoch = load_ckpt_if_exists(gen_state, disc_state)\n\n# def generate_images_jax_single(gen_state_host, example_input, example_target, title_suffix=\"\"):\n#     gen = GeneratorUNet()\n#     vars_ = {'params': gen_state_host.params, 'batch_stats': gen_state_host.batch_stats}\n#     pred = gen.apply(vars_, example_input, train=False, mutable=False)\n\n#     inp = example_input[0]\n#     tar = example_target[0]\n#     out = pred[0]\n\n#     display_input = np.zeros((IMG_HEIGHT, IMG_WIDTH, 3), dtype=np.float32)\n#     display_input[:, :, :2] = np.array(inp)\n\n#     display_list = [display_input, np.array(tar), np.array(out)]\n#     # titles = ['Input (R,G)', 'Ground Truth', f'Predicted {title_suffix}']\n#     titles = ['Label (input. RG)', 'Mask (ground truth)', f'Predicted {title_suffix}']\n#     plt.figure(figsize=(12, 6))\n#     for i in range(3):\n#         plt.subplot(1, 3, i+1)\n#         plt.title(titles[i])\n#         plt.imshow(display_list[i] * 0.5 + 0.5)\n#         plt.axis('off')\n#     plt.show()\n\n# # =========================\n# # 8. TRAIN LOOP\n# # =========================\n# for epoch in range(start_epoch, EPOCHS):\n#     start = time.time()\n#     g_losses, d_losses = [], []\n\n#     for batch_in, batch_tar in train_ds:\n#         batch_in_np, batch_tar_np = tf_batch_to_numpy(batch_in, batch_tar)\n\n#         # enforce correct ordering: inputs = left RG, targets = right RGB\n#         inputs = batch_in_np\n#         targets = batch_tar_np\n\n#         inputs = shard_batch(inputs)\n#         targets = shard_batch(targets)\n\n#         rng, step_rng = random.split(rng)\n#         step_rngs = random.split(step_rng, NUM_DEVICES)\n\n#         gen_state, disc_state, g_loss, d_loss, gan_l, l1_l = p_train_step(\n#             gen_state, disc_state, (inputs, targets), step_rngs\n#         )\n\n#         g_losses.append(np.array(g_loss).mean())\n#         d_losses.append(np.array(d_loss).mean())\n\n#     g_mean = float(np.mean(g_losses))\n#     d_mean = float(np.mean(d_losses))\n#     history['gen_loss'].append(g_mean)\n#     history['disc_loss'].append(d_mean)\n#     print(f\"Epoch {epoch+1}/{EPOCHS} - G: {g_mean:.4f} - D: {d_mean:.4f} - Time: {time.time()-start:.1f}s\")\n\n#     save_ckpt(epoch+1, gen_state, disc_state)\n\n#     gen_state_host = jax.tree_util.tree_map(lambda x: x[0], gen_state)\n#     for example_input, example_target in train_ds.take(1):\n#         ex_in, ex_tar = tf_batch_to_numpy(example_input, example_target)\n#         generate_images_jax_single(gen_state_host, ex_in, ex_tar, f\"Epoch {epoch+1}\")\n#         break\n\n# plt.figure(figsize=(10, 5))\n# plt.plot(history['gen_loss'], label='Gen Loss')\n# plt.plot(history['disc_loss'], label='Disc Loss')\n# plt.title('Loss History')\n# plt.legend()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-24T11:13:33.538640Z","iopub.status.idle":"2026-01-24T11:13:33.538863Z","shell.execute_reply.started":"2026-01-24T11:13:33.538745Z","shell.execute_reply":"2026-01-24T11:13:33.538755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}