{"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":14810989,"sourceType":"datasetVersion","datasetId":9471057},{"sourceId":14829308,"sourceType":"datasetVersion","datasetId":9484229},{"sourceId":290917305,"sourceType":"kernelVersion"},{"sourceId":296912426,"sourceType":"kernelVersion"},{"sourceId":297111656,"sourceType":"kernelVersion"},{"sourceId":297302543,"sourceType":"kernelVersion"},{"sourceId":297674933,"sourceType":"kernelVersion"}],"dockerImageVersionId":31261,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Dependencies & Setup","metadata":{}},{"cell_type":"code","source":"# from IPython.display import clear_output\n\n# var=\"/kaggle/input/notebooks/ipythonx/vsdetection-packages-offline-installer-only/whls\"\n\n# # !pip install \\\n# #   \"$var\"/keras_nightly-*.whl \\\n# #   \"$var\"/tifffile-*.whl \\\n# #   \"$var\"/imagecodecs-*.whl \\\n# #   \"$var\"/medicai-*.whl \\\n# #   --no-index \\\n# #   --find-links \"$var\"\n# var=\"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n# # clear_output()\n# !pip install \\\n#   \"$var\"/tifffile-*.whl \\\n#   \"$var\"/imagecodecs-*.whl \\\n#   \"$var\"/scikit_image-*.whl \\\n#   --no-index \\\n#   --find-links \"$var\n# # clear_output()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T22:15:29.891234Z","iopub.execute_input":"2026-02-15T22:15:29.891375Z","iopub.status.idle":"2026-02-15T22:15:29.898549Z","shell.execute_reply.started":"2026-02-15T22:15:29.891357Z","shell.execute_reply":"2026-02-15T22:15:29.897842Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import clear_output\nvar = \"/kaggle/input/notebooks/crischir/vsdetection-packages-offline-installer-only/whls\"\n\n!pip install \\\n  {var}/tifffile-*.whl \\\n  {var}/imagecodecs-*.whl \\\n  {var}/scikit_image-*.whl \\\n  --no-index \\\n  --find-links {var}\nclear_output()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T22:15:29.898984Z","iopub.execute_input":"2026-02-15T22:15:29.899156Z","iopub.status.idle":"2026-02-15T22:15:37.543579Z","shell.execute_reply.started":"2026-02-15T22:15:29.899139Z","shell.execute_reply":"2026-02-15T22:15:37.542629Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport gc\nimport time\nimport random\nimport warnings\nimport zipfile\nfrom glob import glob\nfrom tqdm.auto import tqdm\n\n# Specialized medical/scientific image formats\nimport tifffile\nimport imagecodecs\n\n# Suppress specific library internal warnings\nwarnings.filterwarnings(\"ignore\", category=FutureWarning, module=\"keras.src.export.tf2onnx_lib\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom PIL import Image, ImageDraw, ImageSequence  # Added ImageSequence for tiff_to_npy\n\nimport jax\nimport jax.numpy as jnp\nfrom jax import grad, jit, vmap, pmap\nimport flax\nfrom flax import linen as nn\nfrom flax.training import train_state, checkpoints\n\n\nfrom flax.core import freeze, unfreeze\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nimport optax\n# import orbax.checkpoint # Uncomment if using Orbax specifically\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom skimage.measure import label, regionprops\nfrom pathlib import Path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T22:15:37.544052Z","iopub.execute_input":"2026-02-15T22:15:37.544249Z","iopub.status.idle":"2026-02-15T22:16:14.074222Z","shell.execute_reply.started":"2026-02-15T22:15:37.544230Z","shell.execute_reply":"2026-02-15T22:16:14.073042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n# MANDATORY: TPU initialization must be at the very top\nos.environ[\"JAX_PLATFORMS\"] = \"tpu\"\n\nimport re\nimport gc\nimport time\nimport random\nimport warnings\nimport zipfile\nfrom glob import glob\nfrom tqdm.auto import tqdm\nfrom pathlib import Path\n\n# Specialized medical/scientific image formats\nimport tifffile\nimport imagecodecs\n\n# Suppress specific library internal warnings\nwarnings.filterwarnings(\"ignore\", category=FutureWarning, module=\"keras.src.export.tf2onnx_lib\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom PIL import Image, ImageDraw, ImageSequence\n\nimport jax\nimport jax.numpy as jnp\nfrom jax import grad, jit, vmap, pmap\nimport flax\nfrom flax import linen as nn\nfrom flax.training import train_state, checkpoints\nfrom flax.core import freeze, unfreeze\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nimport optax\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom skimage.measure import label, regionprops\n\n# Initialize Distributed JAX\ntry:\n    jax.distributed.initialize()\n    print(\"JAX Distributed initialized.\")\nexcept Exception as e:\n    print(f\"Distributed init skipped/failed: {e}\")\n\nprint(f\"Device count: {jax.device_count()}\")\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\nDATA_PATH = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nOUT_DIR = Path(\"/kaggle/temp/vesuvius_npy\")\nOUT_DIR.mkdir(exist_ok=True, parents=True)\nCHECKPOINT_DIR = Path(\"/kaggle/working/checkpoints\")\nCHECKPOINT_DIR.mkdir(parents=True, exist_ok=True)\n\n# Training Hyperparameters\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nBATCH_SIZE = 192 \nLEARNING_RATE_G = 2e-4\nLEARNING_RATE_D = 1e-6\n\n# Increased samples to ensure we have enough data for the 512 BATCH_SIZE\nSAMPLES_PER_ID = 32 \nMAX_EPOCHS = 200\nLOADWeights = True\nMIN_CC_SIZE = 150\nMAX_TRAIN_TIME_HOURS = 3\nMAX_TRAIN_SECONDS = MAX_TRAIN_TIME_HOURS * 3600\n\n# ==========================================\n# 2. UTILS & SCHEDULER\n# ==========================================\ndef get_linear_schedule(init_lr, total_epochs, steps_per_epoch, constant_epochs=30):\n    constant_steps = constant_epochs * steps_per_epoch\n    total_steps = total_epochs * steps_per_epoch\n    decay_steps = max(total_steps - constant_steps, 1)\n    \n    schedule = optax.join_schedules(\n        schedules=[\n            optax.constant_schedule(init_lr),\n            optax.linear_schedule(init_lr, 0.0, decay_steps)\n        ],\n        boundaries=[constant_steps]\n    )\n    return schedule\n\n# ==========================================\n# 3. WEIGHT LOADING (Ignore Mask Source)\n# ==========================================\npath_to_weights = \"/kaggle/input/notebooks/crischir/test-code-for-jax-pix2pix/checkpoints/pix2pix_best.npz\"\nloaded_params, loaded_stats = None, None\n\nif LOADWeights and os.path.exists(path_to_weights):\n    print(f\"Loading weights from {path_to_weights}\")\n    with np.load(path_to_weights) as data:\n        flat_dict = {k: v for k, v in data.items()}\n    \n    p_flat = {k.replace('params/', ''): v for k, v in flat_dict.items() if k.startswith('params/')}\n    s_flat = {k.replace('stats/', ''): v for k, v in flat_dict.items() if k.startswith('stats/')}\n    \n    loaded_params = unflatten_dict({tuple(k.split('/')): v for k, v in p_flat.items()})\n    loaded_stats = unflatten_dict({tuple(k.split('/')): v for k, v in s_flat.items()})\n\n# ==========================================\n# 4. PREPROCESSING\n# ==========================================\ndef tiff_to_npy(tif_path: Path, out_path: Path):\n    if out_path.exists(): return\n    with Image.open(str(tif_path)) as img:\n        frames = [np.array(frame) for frame in ImageSequence.Iterator(img)]\n        vol = np.stack(frames, axis=0)\n    np.save(out_path, vol)\n\ndef run_preprocessing():\n    IMG_SRC = DATA_PATH / \"train_images\"\n    LBL_SRC = DATA_PATH / \"train_labels\"\n    if not IMG_SRC.exists(): return\n    for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\")), desc=\"Images\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n    for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\")), desc=\"Labels\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n    # for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\"))[:350], desc=\"Images\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n    # for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\"))[:350], desc=\"Labels\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n\n# ==========================================\n# 5. DATASET\n# ==========================================\ndef normalize_patch(patch: np.ndarray) -> np.ndarray:\n    patch = patch.astype(np.float32)\n    return (patch - patch.mean()) / (patch.std() + 1e-6)\n\nclass VesuviusSurfaceDataset(Dataset):\n    def __init__(self, ids, npy_dir, samples_per_id, train=True):\n        self.patch_size, self.z_context = PATCH_SIZE, Z_CONTEXT\n        self.rng = np.random.default_rng(42)\n        self.coords, self.volumes, self.labels = [], {}, {}\n        self.train = train \n\n        for vid in ids:\n            img_npy, lbl_npy = npy_dir / f\"{vid}_img.npy\", npy_dir / f\"{vid}_lbl.npy\"\n            if not img_npy.exists() or not lbl_npy.exists(): continue\n            img_vol = np.load(img_npy, mmap_mode=\"r\")\n            z_max, h, w = img_vol.shape\n            z_min, z_max_r = max(self.z_context, 5), min(z_max - self.z_context - 1, z_max - 5)\n            if z_max_r <= z_min: continue\n            zs = self.rng.integers(z_min, z_max_r, size=samples_per_id)\n            ys = self.rng.integers(0, max(1, h - self.patch_size), size=samples_per_id)\n            xs = self.rng.integers(0, max(1, w - self.patch_size), size=samples_per_id)\n            for z, y, x in zip(zs, ys, xs):\n                self.coords.append({\"vid\": vid, \"z\": int(z), \"y\": int(y), \"x\": int(x), \"img_npy\": img_npy, \"lbl_npy\": lbl_npy})\n\n    def __len__(self): return len(self.coords)\n\n    def __getitem__(self, idx):\n        item = self.coords[idx]\n        vid = item[\"vid\"]\n        if vid not in self.volumes:\n            self.volumes[vid] = np.load(item[\"img_npy\"], mmap_mode=\"r\")\n            self.labels[vid] = np.load(item[\"lbl_npy\"], mmap_mode=\"r\")\n        z, y, x = item[\"z\"], item[\"y\"], item[\"x\"]\n        img_patch = self.volumes[vid][z-Z_CONTEXT:z+Z_CONTEXT+1, y:y+PATCH_SIZE, x:x+PATCH_SIZE].transpose(1, 2, 0)\n        z_lbl = z if self.labels[vid].shape[0] > 1 else 0\n        target_patch = self.labels[vid][z_lbl, y:y+PATCH_SIZE, x:x+PATCH_SIZE][..., None]\n        return normalize_patch(img_patch), target_patch.astype(np.float32)\n\n# ==========================================\n# 6. MODELS\n# ==========================================\nclass Pix2PixGenerator(nn.Module):\n    @nn.compact\n    def __call__(self, x, train: bool = True):\n        def conv_bn_leaky(feat, out_c):\n            feat = nn.Conv(out_c, (4,4), strides=(2,2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            return nn.leaky_relu(feat, 0.2)\n        \n        s1 = nn.leaky_relu(nn.Conv(64, (4,4), strides=(2,2), padding='SAME')(x), 0.2)\n        s2 = conv_bn_leaky(s1, 128)\n        s3 = conv_bn_leaky(s2, 256)\n        \n        def up_bn_relu(feat, skip, out_c):\n            feat = nn.ConvTranspose(out_c, (4,4), strides=(2,2), padding='SAME')(feat)\n            feat = nn.BatchNorm(use_running_average=not train)(feat)\n            feat = nn.relu(feat)\n            return jnp.concatenate([feat, skip], axis=-1)\n        \n        u1 = up_bn_relu(s3, s2, 128)\n        u2 = up_bn_relu(u1, s1, 64)\n        return nn.ConvTranspose(1, (4,4), strides=(2,2), padding='SAME')(u2)\n\nclass TopologyDiscriminator(nn.Module):\n    @nn.compact\n    def __call__(self, x, label, train: bool = True):\n        inputs = jnp.concatenate([x, label], axis=-1)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(64, (4,4), strides=(2,2), padding='SAME')(inputs)), 0.2)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(128, (4,4), strides=(2,2), padding='SAME')(y)), 0.2)\n        y = nn.leaky_relu(nn.BatchNorm(use_running_average=not train)(nn.Conv(256, (4,4), strides=(2,2), padding='SAME')(y)), 0.2)\n        return nn.Conv(1, (4,4), padding='SAME')(y)\n\nclass JaxTrainState(train_state.TrainState):\n    batch_stats: flax.core.FrozenDict\n\ndef save_portable_npy(state, filename):\n    flat_params = flatten_dict(jax.device_get(state.params), sep='/')\n    flat_stats = flatten_dict(jax.device_get(state.batch_stats), sep='/')\n    save_dict = {f\"params/{k}\": v for k, v in flat_params.items()}\n    save_dict.update({f\"stats/{k}\": v for k, v in flat_stats.items()})\n    np.savez(CHECKPOINT_DIR / filename, **save_dict)\n    print(f\"⭐ Portable .npz weights saved: {filename}\")\n\n# ==========================================\n# 7. LOSS & TRAINING (With Ignore Mask)\n# ==========================================\ndef hybrid_loss(logits, target, ignore_mask):\n    valid_mask = (target != 2).astype(jnp.float32)\n    # Mask is 0 where pre-trained confidence > 0.7, 1 elsewhere\n    combined_mask = valid_mask * ignore_mask\n    \n    target_binary = (target == 1).astype(jnp.float32)\n    bce = optax.sigmoid_binary_cross_entropy(logits, target_binary)\n    \n    mask_sum = jnp.sum(combined_mask)\n    # Use jnp.where to prevent NaN if the entire batch is masked\n    masked_bce = jnp.where(mask_sum > 0, jnp.sum(bce * combined_mask) / (mask_sum + 1e-8), 0.0)\n    \n    pred = nn.sigmoid(logits)\n    dice_num = 2.0 * jnp.sum(pred * target_binary * combined_mask)\n    dice_den = jnp.sum(pred * combined_mask) + jnp.sum(target_binary * combined_mask)\n    dice = 1 - (dice_num + 1e-8) / (dice_den + 1e-8)\n    \n    return 0.5 * masked_bce + 0.5 * dice\n\n@jax.jit\ndef train_step(state_g, state_d, state_ignore, imgs, targets, d_train_flag):\n    # Step 0: Generate Ignore Mask from loaded (frozen) weights\n    ignore_out = state_ignore.apply_fn(\n        {'params': state_ignore.params, 'batch_stats': state_ignore.batch_stats}, \n        imgs, train=False\n    )\n    # If prediction > 0.7, mask value = 0 (ignore)\n    ignore_mask = (nn.sigmoid(ignore_out) <= 0.7).astype(jnp.float32)\n\n    def loss_d_fn(params):\n        fake_targets = nn.sigmoid(state_g.apply_fn({'params': state_g.params, 'batch_stats': state_g.batch_stats}, imgs, train=False))\n        combined_imgs = jnp.concatenate([imgs, imgs], axis=0)\n        combined_lbls = jnp.concatenate([(targets == 1).astype(jnp.float32), fake_targets], axis=0)\n        (d_out, updates) = state_d.apply_fn({'params': params, 'batch_stats': state_d.batch_stats}, combined_imgs, combined_lbls, train=True, mutable=['batch_stats'])\n        d_real, d_fake = jnp.split(d_out, 2, axis=0)\n        return jnp.mean((d_real - 1)**2) + jnp.mean(d_fake**2), updates\n\n    def loss_g_fn(params):\n        (gen_logits, updates) = state_g.apply_fn({'params': params, 'batch_stats': state_g.batch_stats}, imgs, train=True, mutable=['batch_stats'])\n        d_out = state_d.apply_fn({'params': state_d.params, 'batch_stats': state_d.batch_stats}, imgs, nn.sigmoid(gen_logits), train=False)\n        return hybrid_loss(gen_logits, targets, ignore_mask) + 0.1 * jnp.mean((d_out - 1)**2), updates\n\n    # Optimize D\n    grad_d_fn = jax.value_and_grad(loss_d_fn, has_aux=True)\n    (loss_d, d_stats), grad_d = grad_d_fn(state_d.params)\n    state_d = jax.lax.cond(d_train_flag, lambda s,g,st: s.apply_gradients(grads=g).replace(batch_stats=st['batch_stats']), lambda s,g,st: s, state_d, grad_d, d_stats)\n\n    # Optimize G\n    grad_g_fn = jax.value_and_grad(loss_g_fn, has_aux=True)\n    (loss_g, g_stats), grad_g = grad_g_fn(state_g.params)\n    state_g = state_g.apply_gradients(grads=grad_g).replace(batch_stats=g_stats['batch_stats'])\n    \n    return state_g, state_d, loss_g, loss_d\n\n# ==========================================\n# 8. VISUALIZATION\n# ==========================================\ndef apply_cc_filter(mask, min_size=MIN_CC_SIZE):\n    mask = (mask > 0.25).astype(np.uint8)\n    labels = label(mask)\n    for region in regionprops(labels):\n        if region.area < min_size:\n            mask[labels == region.label] = 0\n    return mask\n\ndef visualize_results(state_g, dataset, epoch):\n    rows, cols = 5, 7\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4))\n    for i in range(rows):\n        img_np, target_np = dataset[random.randint(0, len(dataset)-1)]\n        logits = state_g.apply_fn({'params': state_g.params, 'batch_stats': state_g.batch_stats}, jnp.array(img_np[None, ...]), train=False)\n        pred = np.array(nn.sigmoid(logits)).squeeze()\n        thr = (pred > 0.5).astype(np.uint8)\n        cleaned = apply_cc_filter(thr.copy(), 150)\n        target_viz = np.zeros((*target_np.shape[:2], 3))\n        target_viz[target_np.squeeze() == 1] = [1, 1, 1]\n        target_viz[target_np.squeeze() == 2] = [0.5, 0.5, 0.5]\n        overlay_in = np.zeros((*cleaned.shape, 3)); overlay_in[..., 2] = 0.3; overlay_in[..., 0] = cleaned\n        overlay_lbl = target_viz.copy(); overlay_lbl[..., 0] = np.maximum(overlay_lbl[..., 0], cleaned.astype(float))\n        axes[i, 0].imshow(img_np[..., Z_CONTEXT], cmap=\"gray\"); axes[i, 1].imshow(target_viz)\n        axes[i, 2].imshow(pred, cmap=\"magma\"); axes[i, 3].imshow(thr, cmap=\"gray\")\n        axes[i, 4].imshow(cleaned, cmap=\"gray\"); axes[i, 5].imshow(overlay_in); axes[i, 6].imshow(overlay_lbl)\n        for j in range(cols): axes[i, j].axis(\"off\")\n    plt.tight_layout(); plt.savefig(f\"viz_{epoch}.png\"); plt.show(); plt.close()\n\n# ==========================================\n# 9. MAIN EXECUTION\n# ==========================================\nif __name__ == \"__main__\":\n    run_preprocessing()\n    ids = sorted({p.stem.replace(\"_img\", \"\") for p in OUT_DIR.glob(\"*_img.npy\")})\n    ds = VesuviusSurfaceDataset(ids, OUT_DIR, SAMPLES_PER_ID, train=True)\n    \n    # drop_last=False ensures we process even small datasets\n    loader = DataLoader(ds, batch_size=BATCH_SIZE, shuffle=True, drop_last=False, num_workers=0)\n    steps_per_epoch = len(loader)\n    print(f\"Dataset samples: {len(ds)} | Steps per epoch: {steps_per_epoch}\")\n\n    rng = jax.random.PRNGKey(42)\n    netG, netD = Pix2PixGenerator(), TopologyDiscriminator()\n    dummy = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 2*Z_CONTEXT+1))\n\n    sched_G = get_linear_schedule(LEARNING_RATE_G, MAX_EPOCHS, steps_per_epoch)\n    sched_D = get_linear_schedule(LEARNING_RATE_D, MAX_EPOCHS, steps_per_epoch)\n    \n    # Initialize Generator\n    v_g = netG.init(jax.random.split(rng)[0], dummy)\n    \n    if LOADWeights and loaded_params is not None:\n        state_g = JaxTrainState.create(\n            apply_fn=netG.apply, params=loaded_params, \n            batch_stats=loaded_stats, tx=optax.adam(sched_G)\n        )\n        # Frozen state for the Ignore Mask prediction\n        state_ignore = JaxTrainState.create(\n            apply_fn=netG.apply, params=loaded_params, \n            batch_stats=loaded_stats, tx=optax.identity()\n        )\n    else:\n        state_g = JaxTrainState.create(\n            apply_fn=netG.apply, params=v_g['params'], \n            batch_stats=v_g['batch_stats'], tx=optax.adam(sched_G, 0.5, 0.999)\n        )\n        state_ignore = state_g # Fallback if no weights found\n\n    # Initialize Discriminator\n    v_d = netD.init(jax.random.split(rng)[1], dummy, jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1)))\n    state_d = JaxTrainState.create(\n        apply_fn=netD.apply, params=v_d['params'], \n        batch_stats=v_d['batch_stats'], tx=optax.adam(sched_D, 0.5, 0.999)\n    )\n\n    best_loss = float('inf')\n    start_time = time.time() \n    \n    for ep in range(1, MAX_EPOCHS + 1):\n        if (time.time() - start_time) > MAX_TRAIN_SECONDS:\n            print(\"⏳ Max training time reached.\")\n            break\n        \n        epoch_g_loss, epoch_d_loss = [], []\n        \n        for step, (b_imgs, b_tgts) in enumerate(loader):\n            # Pass state_ignore to use it for masking\n            state_g, state_d, lg, ld = train_step(\n                state_g, state_d, state_ignore, jnp.array(b_imgs), jnp.array(b_tgts), step % 2 == 0\n            )\n            epoch_g_loss.append(lg)\n            epoch_d_loss.append(ld)\n\n            if step % 20 == 0:\n                print(f\"Ep {ep} | Step {step}/{steps_per_epoch} | G_loss: {lg:.4f} | D_loss: {ld:.4f}\")\n\n        if len(epoch_g_loss) == 0: continue\n\n        avg_g, avg_d = np.mean(epoch_g_loss), np.mean(epoch_d_loss)\n        elapsed = (time.time() - start_time) / 3600\n        print(f\"✅ Epoch {ep} Complete | Avg G: {avg_g:.4f} | Avg D: {avg_d:.4f} | Time: {elapsed:.2f}h\")\n\n        if ep % 10 == 0:\n            visualize_results(state_g, ds, ep)\n            save_portable_npy(state_g, f\"pix2pix_ep_{ep}.npz\")\n            \n        if avg_g < best_loss:\n            best_loss = float(avg_g)\n            save_portable_npy(state_g, \"pix2pix_best.npz\")\n\n    # ==========================================\n    # 10. INFERENCE & SUBMISSION\n    # ==========================================\n    test_dir = DATA_PATH / 'test_images'\n    if test_dir.exists():\n        for vid in [f.stem for f in test_dir.glob('*.tif')]:\n            with Image.open(str(test_dir / f\"{vid}.tif\")) as img:\n                vol = np.stack([np.array(f) for f in ImageSequence.Iterator(img)], axis=0)\n            z_m, h, w = vol.shape\n            pages = []\n            for z in tqdm(range(z_m), desc=f\"Infer {vid}\"):\n                if z-Z_CONTEXT < 0 or z+Z_CONTEXT+1 > z_m:\n                    pages.append(Image.fromarray(np.zeros((h, w), dtype=np.uint8)))\n                    continue\n                pp = np.zeros((h, w))\n                for y in range(0, h, PATCH_SIZE):\n                    for x in range(0, w, PATCH_SIZE):\n                        y1, x1 = min(y+PATCH_SIZE, h), min(x+PATCH_SIZE, w)\n                        y0, x0 = y1-PATCH_SIZE, x1-PATCH_SIZE\n                        p = normalize_patch(vol[z-Z_CONTEXT:z+Z_CONTEXT+1, y0:y1, x0:x1].transpose(1, 2, 0))\n                        o = state_g.apply_fn({'params': state_g.params, 'batch_stats': state_g.batch_stats}, jnp.array(p[None, ...]), train=False)\n                        pp[y0:y1, x0:x1] = np.array(nn.sigmoid(o)).squeeze()\n                pages.append(Image.fromarray(apply_cc_filter((pp > 0.2).astype(np.uint8), 150)))\n            pages[0].save(f\"{vid}.tif\", save_all=True, append_images=pages[1:], compression=\"tiff_deflate\")\n        \n        with zipfile.ZipFile('submission.zip', 'w') as z:\n            for f in Path('.').glob('*.tif'):\n                z.write(f)\n                os.remove(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T22:31:14.149890Z","iopub.execute_input":"2026-02-15T22:31:14.150243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"my_path=\"/kaggle/working/submission.zip\"\ndef visualize_and_verify_zip(zip_path=my_path, samples_per_volume=15):\n    \"\"\"\n    Independent verification script for Vesuvius Challenge submissions.\n    Checks for binary integrity and visualizes start, middle, and end slices.\n    \"\"\"\n    if not os.path.exists(zip_path):\n        print(f\"❌ Error: {zip_path} not found in current directory.\")\n        return\n\n    extract_dir = 'verification_extract'\n    os.makedirs(extract_dir, exist_ok=True)\n\n    with zipfile.ZipFile(zip_path, 'r') as zip_ref:\n        zip_ref.extractall(extract_dir)\n        tif_files = [f for f in zip_ref.namelist() if f.endswith('.tif')]\n\n    if not tif_files:\n        print(\"❌ No .tif files found inside the zip.\")\n        return\n\n    print(f\"📦 Found {len(tif_files)} volumes in submission.zip\\n\")\n\n    for tif_file in tif_files:\n        file_path = os.path.join(extract_dir, tif_file)\n        \n        # 1. Load the volume\n        volume = tifffile.imread(file_path)\n        z_max, h, w = volume.shape\n\n        # 2. Binary Integrity Check\n        unique_vals = np.unique(volume)\n        is_binary = np.all(np.isin(unique_vals, [0, 1]))\n        \n        print(f\"--- Volume: {tif_file} ---\")\n        print(f\"Dimensions: {z_max} layers | Resolution: {w}x{h}\")\n        print(f\"Unique values in file: {unique_vals}\")\n        \n        if is_binary:\n            print(\"✅ Status: Correctly formatted (Binary 0/1)\")\n        else:\n            print(\"⚠️ Status: NOT BINARY (Check your thresholding/saving logic)\")\n\n        # 3. Sampling Logic (Start, Middle, End)\n        # Grab 5 indices from the first 50, 5 from the middle, 5 from the last 50\n        first_idx = random.sample(range(0, min(50, z_max)), 5)\n        mid_idx = random.sample(range(max(0, z_max//2 - 25), min(z_max, z_max//2 + 25)), 5)\n        last_idx = random.sample(range(max(0, z_max - 50), z_max), 5)\n        \n        all_indices = sorted(first_idx + mid_idx + last_idx)\n\n        # 4. Visualization Grid\n        fig, axes = plt.subplots(3, 5, figsize=(20, 12))\n        fig.suptitle(f\"Verification: {tif_file} (Z-Slices sampled from Start, Middle, End)\", fontsize=18)\n\n        for i, idx in enumerate(all_indices):\n            row = i // 5\n            col = i % 5\n            \n            slice_data = volume[idx]\n            \n            # Using 'gray' cmap: 0 = Black, 1 = White\n            axes[row, col].imshow(slice_data, cmap='gray', interpolation='nearest')\n            axes[row, col].set_title(f\"Z-Index: {idx}\")\n            axes[row, col].axis('off')\n\n            # Check if slice is empty\n            if np.sum(slice_data) == 0:\n                axes[row, col].text(w//2, h//2, \"EMPTY\", color='red', \n                                   ha='center', va='center', fontweight='bold')\n\n        plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n        plt.show()\n\n    # Optional: Clean up extracted files to save space\n    # import shutil\n    # shutil.rmtree(extract_dir)\n\nif __name__ == \"__main__\":\n    visualize_and_verify_zip()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T22:17:40.186879Z","iopub.status.idle":"2026-02-15T22:17:40.187324Z","shell.execute_reply.started":"2026-02-15T22:17:40.186993Z","shell.execute_reply":"2026-02-15T22:17:40.187004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}