{"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"}],"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-14T11:22:45.427305Z","iopub.execute_input":"2026-02-14T11:22:45.427454Z","iopub.status.idle":"2026-02-14T11:22:45.441576Z","shell.execute_reply.started":"2026-02-14T11:22:45.427436Z","shell.execute_reply":"2026-02-14T11:22:45.440880Z"},"_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-14T11:22:45.442169Z","iopub.execute_input":"2026-02-14T11:22:45.442315Z","iopub.status.idle":"2026-02-14T11:22:51.017027Z","shell.execute_reply.started":"2026-02-14T11:22:45.442299Z","shell.execute_reply":"2026-02-14T11:22:51.016196Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from flax.training import train_state, checkpoints\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nfrom scipy import ndimage\nfrom skimage.measure import label, regionprops\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=FutureWarning, module=\"keras.src.export.tf2onnx_lib\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-14T11:22:51.017539Z","iopub.execute_input":"2026-02-14T11:22:51.017712Z","iopub.status.idle":"2026-02-14T11:23:29.386797Z","shell.execute_reply.started":"2026-02-14T11:22:51.017695Z","shell.execute_reply":"2026-02-14T11:23:29.385773Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import imagecodecs\nimport tifffile","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-14T11:23:29.387618Z","iopub.execute_input":"2026-02-14T11:23:29.388048Z","iopub.status.idle":"2026-02-14T11:23:29.399531Z","shell.execute_reply.started":"2026-02-14T11:23:29.388029Z","shell.execute_reply":"2026-02-14T11:23:29.398743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n# !pip install -q imagecodecs flax optax scikit-image\nimport os\nimport time\nimport random\nimport zipfile\nimport numpy as np\n\nimport multiprocessing\n# 2. Unset Kaggle proxy variables that cause the distributed hang/crash\n# for var in [\"KAGGLE_DATA_PROXY_TOKEN\", \"KAGGLE_DATA_PROXY_PROJECT\", \n#             \"KAGGLE_GRPC_DATA_PROXY_URL\", \"KAGGLE_DATA_PROXY_URL\"]:\n#     os.environ.pop(var, None)\n\n# 3. Force JAX to use the local TPU VM backend\nos.environ[\"JAX_PLATFORMS\"] = \"tpu\"\n\nimport jax\n# 4. Initialize without a coordinator address to stick to local TPU VM\ntry:\n    jax.distributed.initialize()\n    print(\"JAX Distributed initialized.\")\nexcept Exception as e:\n    print(f\"Standard init failed, trying local: {e}\")\n\nprint(f\"Device count: {jax.device_count()}\")\n\nimport jax.numpy as jnp\nimport flax\nimport flax.linen as nn\nimport optax\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom PIL import Image, ImageSequence\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\nfrom flax.training import train_state, checkpoints\nfrom flax.traverse_util import flatten_dict, unflatten_dict\nfrom scipy import ndimage\nfrom skimage.measure import label, regionprops\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\n# !pip install imagecodecs flax optax\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 (Preserved)\nPATCH_SIZE = 256\nZ_CONTEXT = 3\nBATCH_SIZE = 256 \nLEARNING_RATE_G = 2e-4\nLEARNING_RATE_D = 1e-6\n\nSAMPLES_PER_ID = 64\nMAX_BATCHES_PER_EPOCH = 1800\nMAX_TRAIN_TIME_HOURS = 1.5\nMAX_TRAIN_SECONDS = MAX_TRAIN_TIME_HOURS * 3600\nMAX_EPOCHS = 200\nLOADWeights=True\n# ==========================================\n# SCHEDULER HELPER\n# ==========================================\ndef get_linear_schedule(init_lr, total_epochs, steps_per_epoch, constant_epochs=20):\n    \"\"\"Creates a linear decay schedule after an initial constant phase.\"\"\"\n    constant_steps = constant_epochs * steps_per_epoch\n    total_steps = total_epochs * steps_per_epoch\n    decay_steps = total_steps - constant_steps\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# 1. Load existing weights\n# ==========================================\n# path_to_weights=\"/kaggle/input/notebooks/crischir/jax-pix2pix/checkpoints/pix2pix_best.npz\"\npath_to_weights=\"/kaggle/input/datasets/crischir/vesuvius-checkpoint-dataset/checkpoints/pix2pix_ep_320.npz\"\nif LOADWeights:\n    with np.load(path_to_weights) as data:\n        flat_dict = {k: v for k, v in data.items()}\n    \n    # 2. Separate and unflatten\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\n# ==========================================\n# 2. PREPROCESSING (TIFF to NPY)\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    for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\")), desc=\"Preprocessing Images\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n    for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\")), desc=\"Preprocessing Labels\"):\n        tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n\n    # for tif in tqdm(sorted(IMG_SRC.glob(\"*.tif\"))[:10], desc=\"Images\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_img.npy\")\n\n    # # Slicing the first 10 labels\n    # for tif in tqdm(sorted(LBL_SRC.glob(\"*.tif\"))[:10], desc=\"Labels\"):\n    #     tiff_to_npy(tif, OUT_DIR / f\"{tif.stem}_lbl.npy\")\n\n# ==========================================\n# 3. DATASET & AUGMENTATION (TTA Logic)\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 # Toggles Augmentation\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 _random_augment(self, img, target):\n        \"\"\" Ported exactly from original logic, adjusted for HWC \"\"\"\n        if self.rng.random() > 0.5: # Horizontal Flip\n            img, target = np.flip(img, axis=1), np.flip(target, axis=1)\n        if self.rng.random() > 0.5: # Vertical Flip\n            img, target = np.flip(img, axis=0), np.flip(target, axis=0)\n        k = self.rng.integers(0, 4)\n        if k > 0: # 90-degree Rotations\n            img, target = np.rot90(img, k=k, axes=(0, 1)), np.rot90(target, k=k, axes=(0, 1))\n        return img.copy(), target.copy()\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        \n        # 2.5D Sandwich Slicing\n        img_patch = self.volumes[vid][z-Z_CONTEXT:z+Z_CONTEXT+1, y:y+PATCH_SIZE, x:x+PATCH_SIZE]\n        img_patch = img_patch.transpose(1, 2, 0) # to HWC\n        \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        \n        if self.train:\n            img_patch, target_patch = self._random_augment(img_patch, target_patch)\n            \n        return normalize_patch(img_patch), target_patch.astype(np.float32)\n\n    def __len__(self): return len(self.coords)\n\n# ==========================================\n# 4. MODELS (Flax - Channels Last)\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 TrainState(train_state.TrainState):\n    batch_stats: flax.core.FrozenDict\nclass JaxTrainState(train_state.TrainState):\n    batch_stats: flax.core.FrozenDict\ndef save_portable_npy(state, filename):\n    \"\"\" Securely saves weights as a flattened .npz file \"\"\"\n    # Convert nested dict to flat dict with '/' keys\n    flat_params = flatten_dict(jax.device_get(state.params), sep='/')\n    flat_stats = flatten_dict(jax.device_get(state.batch_stats), sep='/')\n    \n    # Prefix keys to distinguish params from stats\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    \n    np.savez(CHECKPOINT_DIR / filename, **save_dict)\n    print(f\"⭐ Portable .npz weights saved: {filename}\")\n# ==========================================\n# 5. LOSS & TRAINING (JIT Compiled)\n# ==========================================\ndef hybrid_loss(logits, target):\n    valid_mask = (target != 2).astype(jnp.float32)\n    target_binary = (target == 1).astype(jnp.float32)\n    bce = optax.sigmoid_binary_cross_entropy(logits, target_binary)\n    masked_bce = jnp.sum(bce * valid_mask) / (jnp.sum(valid_mask) + 1e-6)\n    pred = nn.sigmoid(logits)\n    dice = 1 - (2.0 * jnp.sum(pred * target_binary * valid_mask) + 1e-6) / \\\n           (jnp.sum(pred * valid_mask) + jnp.sum(target_binary * valid_mask) + 1e-6)\n    return 0.5 * masked_bce + 0.5 * dice\n\n@jax.jit\ndef train_step(state_g, state_d, imgs, targets, d_train_flag):\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) + 0.1 * jnp.mean((d_out - 1)**2), updates\n\n    # Step 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    # Step 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# 6. VISUALIZATION & POST-PROCESSING\n# ==========================================\ndef apply_cc_filter(mask: np.ndarray, min_size: int = 150):\n    labeled, num_features = ndimage.label(mask)\n    if num_features == 0: return mask\n    component_sizes = np.bincount(labeled.ravel())\n    mask[component_sizes[labeled] < min_size] = 0\n    return mask\n\ndef visualize_cc_impact(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        \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        \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\n        axes[i, 0].imshow(img_np[..., Z_CONTEXT], cmap=\"gray\"); axes[i, 0].set_title(\"Input Slice\")\n        axes[i, 1].imshow(target_viz); axes[i, 1].set_title(\"Target Label\")\n        axes[i, 2].imshow(pred, cmap=\"magma\"); axes[i, 2].set_title(\"Raw Pred\")\n        axes[i, 3].imshow(thr, cmap=\"gray\"); axes[i, 3].set_title(\"Thresholded\")\n        axes[i, 4].imshow(cleaned, cmap=\"gray\"); axes[i, 4].set_title(\"Cleaned CC\")\n        axes[i, 5].imshow(overlay_in); axes[i, 5].set_title(\"Overlay Input\")\n        axes[i, 6].imshow(overlay_lbl); axes[i, 6].set_title(\"Overlay Label\")\n        for j in range(cols): axes[i, j].axis(\"off\")\n    plt.tight_layout(); plt.savefig(f\"visual_epoch_{epoch}.png\"); plt.show(); plt.close()\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# 7. 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    # loader = DataLoader(ds, batch_size=BATCH_SIZE, shuffle=True, drop_last=True, num_workers=4)\n    loader = DataLoader(ds, batch_size=BATCH_SIZE, shuffle=True, drop_last=True, num_workers=0)\n    steps_per_epoch = len(loader)\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    # Define Schedulers\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    \n    v_g = netG.init(jax.random.split(rng)[0], dummy)\n    #No learning Schedueler\n    # if LOADWeights:\n    #     # 3. Create state with loaded weights\n    #     state_g = JaxTrainState.create(\n    #         apply_fn=netG.apply, \n    #         params=loaded_params, \n    #         batch_stats=loaded_stats, \n    #         tx=optax.adam(LEARNING_RATE_G)\n    #     )\n    # else:\n    #      state_g = JaxTrainState.create(apply_fn=netG.apply, params=v_g['params'], batch_stats=v_g['batch_stats'], tx=optax.adam(LEARNING_RATE_G, 0.5, 0.999))\n\n    if LOADWeights:\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    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\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    \n    v_d = netD.init(jax.random.split(rng)[1], dummy, jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1)))\n    state_d = JaxTrainState.create(apply_fn=netD.apply, params=v_d['params'], batch_stats=v_d['batch_stats'], tx=optax.adam(LEARNING_RATE_D, 0.5, 0.999))\n\n    best_loss, start = float('inf'), time.time()\n    start_time = time.time()  # This was likely named 'start' in your old code\n    # for ep in range(1, MAX_EPOCHS + 1):\n    #     if (time.time() - start) > MAX_TRAIN_SECONDS: break\n    #     pbar = tqdm(loader, desc=f\"Ep {ep}\")\n    #     for step, (b_imgs, b_tgts) in enumerate(pbar):\n    #         state_g, state_d, lg, ld = train_step(state_g, state_d, jnp.array(b_imgs), jnp.array(b_tgts), step % 2 == 0)\n    #         pbar.set_postfix(G_loss=f\"{lg:.4f}\", D_loss=f\"{ld:.4f}\")\n        \n    #     if ep % 10 == 0: visualize_results(state_g, ds, ep)\n        \n    #     # Periodic weight-only portable saving (Every 10 epochs)\n    #     if ep % 10 == 0:\n    #         save_portable_npy(state_g, f\"pix2pix_ep_{ep}.npz\")\n            \n    #     if lg < best_loss:\n    #         best_loss = float(lg)\n    #         save_portable_npy(state_g, \"pix2pix_best.npz\")\n    # print(f\"Starting training for {MAX_EPOCHS} epochs...\")\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        # Clean logging: No tqdm, just a simple loop\n        for step, (b_imgs, b_tgts) in enumerate(loader):\n            state_g, state_d, lg, ld = train_step(\n                state_g, state_d, 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            # Print status every 100 batches instead of a progress bar\n            if step % 100 == 0:\n                print(f\"Ep {ep} | Step {step}/{steps_per_epoch} | G_loss: {lg:.4f} | D_loss: {ld:.4f}\")\n\n        # End of Epoch Summary\n        avg_g = np.mean(epoch_g_loss)\n        avg_d = 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        # Visualization and Checkpointing\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    # Final Submission Inference\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        with zipfile.ZipFile('submission.zip', 'w') as z:\n            for f in Path('.').glob('*.tif'): z.write(f); os.remove(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-14T11:23:29.400421Z","iopub.execute_input":"2026-02-14T11:23:29.400605Z","execution_failed":"2026-02-14T11:26:00.803Z"}},"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":{"execution_failed":"2026-02-14T11:26:00.804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}