{"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"},{"sourceId":297732241,"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-15T00:44:37.389514Z","iopub.execute_input":"2026-02-15T00:44:37.389812Z","iopub.status.idle":"2026-02-15T00:44:37.393637Z","shell.execute_reply.started":"2026-02-15T00:44:37.389789Z","shell.execute_reply":"2026-02-15T00:44:37.392528Z"},"_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-15T00:44:38.100799Z","iopub.execute_input":"2026-02-15T00:44:38.101126Z","iopub.status.idle":"2026-02-15T00:44:44.340443Z","shell.execute_reply.started":"2026-02-15T00:44:38.101099Z","shell.execute_reply":"2026-02-15T00:44:44.339533Z"},"_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-15T00:44:44.341174Z","iopub.execute_input":"2026-02-15T00:44:44.341372Z","iopub.status.idle":"2026-02-15T00:44:44.358258Z","shell.execute_reply.started":"2026-02-15T00:44:44.341353Z","shell.execute_reply":"2026-02-15T00:44:44.357547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import imagecodecs\nimport tifffile","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-15T00:44:44.358648Z","iopub.execute_input":"2026-02-15T00:44:44.359217Z","iopub.status.idle":"2026-02-15T00:44:44.372571Z","shell.execute_reply.started":"2026-02-15T00:44:44.359199Z","shell.execute_reply":"2026-02-15T00:44:44.371909Z"}},"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\nimport shutil\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# --- LOSS CONFIG ---\nL1_LAMBDA = 10.0\nSOBEL_LAMBDA = 1.0\n# -----------------------\n\nSAMPLES_PER_ID = 24\nMAX_BATCHES_PER_EPOCH = 1800\nMAX_TRAIN_TIME_HOURS = 1.2\nMAX_TRAIN_SECONDS = MAX_TRAIN_TIME_HOURS * 3600\nMAX_EPOCHS = 1000\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/notebooks/crischir/test-code-sobel-jax-pix2pix/checkpoints/pix2pix_ep_130.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 get_prewitt_kernels():\n    \"\"\"Defines Prewitt kernels: [-1, -1, -1].\"\"\"\n    # Vertical Kernel (detects horizontal lines)\n    k_y = jnp.array([[-1, -1, -1], [0, 0, 0], [1, 1, 1]], dtype=jnp.float32)\n    # Horizontal Kernel (detects vertical lines)\n    k_x = jnp.array([[-1, 0, 1], [-1, 0, 1], [-1, 0, 1]], dtype=jnp.float32)\n    \n    # Expand dims for JAX: (H, W, In, Out) -> (3, 3, 1, 1)\n    k_y = k_y[:, :, jnp.newaxis, jnp.newaxis]\n    k_x = k_x[:, :, jnp.newaxis, jnp.newaxis]\n    return k_x, k_y\n\ndef compute_masked_prewitt_loss(y_true, y_pred, valid_mask):\n    \"\"\"\n    Calculates Prewitt loss using L1 Norm (Manhattan Distance).\n    Avoids sqrt() to prevent NaN on TPUs.\n    \"\"\"\n    k_x, k_y = get_prewitt_kernels()\n    dn = jax.lax.ConvDimensionNumbers((0, 3, 1, 2), (3, 2, 0, 1), (0, 3, 1, 2))\n    \n    def get_gradients(img):\n        grad_x = jax.lax.conv_general_dilated(img, k_x, (1, 1), 'SAME', dimension_numbers=dn)\n        grad_y = jax.lax.conv_general_dilated(img, k_y, (1, 1), 'SAME', dimension_numbers=dn)\n        return grad_x, grad_y\n\n    # Get gradients for Ground Truth and Prediction\n    gx_true, gy_true = get_gradients(y_true)\n    gx_pred, gy_pred = get_gradients(y_pred)\n    \n    # --- STABILITY FIX ---\n    # Instead of sqrt(gx^2 + gy^2), we use |gx| + |gy|\n    # We compare the X-gradients and Y-gradients separately\n    diff_x = jnp.abs(gx_true - gx_pred)\n    diff_y = jnp.abs(gy_true - gy_pred)\n    \n    # Sum errors and apply mask\n    total_diff = (diff_x + diff_y) * valid_mask\n    \n    return jnp.sum(total_diff) / (jnp.sum(valid_mask) + 1e-6)\n\ndef hybrid_loss(logits, target):\n    # 1. Prepare Masks\n    valid_mask = (target != 2).astype(jnp.float32)\n    target_binary = (target == 1).astype(jnp.float32)\n    \n    # Clip sigmoid to avoid 0 or 1 exactly\n    pred = jnp.clip(nn.sigmoid(logits), 1e-7, 1 - 1e-7)\n\n    # 2. Standard Losses (BCE + Dice)\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    \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    \n    base_loss = 0.5 * masked_bce + 0.5 * dice\n\n    # 3. L1 Pixel Loss\n    l1_diff = jnp.abs(pred - target_binary) * valid_mask\n    l1_loss = jnp.sum(l1_diff) / (jnp.sum(valid_mask) + 1e-6)\n\n    # 4. Prewitt Edge Loss (Stable L1 version)\n    edge_loss = compute_masked_prewitt_loss(target_binary, pred, valid_mask)\n\n    # 5. Combine\n    total_loss = base_loss + (L1_LAMBDA * l1_loss) + (SOBEL_LAMBDA * edge_loss)\n    \n    return total_loss\ndef get_sobel_kernels():\n    \"\"\"Defines Sobel kernels for edge detection in JAX (H, W, In, Out).\"\"\"\n    k_y = jnp.array([[-1, -2, -1], [0, 0, 0], [1, 2, 1]], dtype=jnp.float32)\n    k_x = jnp.array([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], dtype=jnp.float32)\n    # Expand dims to (H, W, In, Out) -> (3, 3, 1, 1)\n    k_y = k_y[:, :, jnp.newaxis, jnp.newaxis]\n    k_x = k_x[:, :, jnp.newaxis, jnp.newaxis]\n    return k_x, k_y\n\ndef compute_masked_prewitt_loss(y_true, y_pred, valid_mask):\n    k_x, k_y = get_prewitt_kernels()\n    dn = jax.lax.ConvDimensionNumbers((0, 3, 1, 2), (3, 2, 0, 1), (0, 3, 1, 2))\n    \n    # Calculate gradients\n    gx_true = jax.lax.conv_general_dilated(y_true, k_x, (1, 1), 'SAME', dimension_numbers=dn)\n    gy_true = jax.lax.conv_general_dilated(y_true, k_y, (1, 1), 'SAME', dimension_numbers=dn)\n    gx_pred = jax.lax.conv_general_dilated(y_pred, k_x, (1, 1), 'SAME', dimension_numbers=dn)\n    gy_pred = jax.lax.conv_general_dilated(y_pred, k_y, (1, 1), 'SAME', dimension_numbers=dn)\n\n    # L1 Difference (Safe)\n    diff = (jnp.abs(gx_true - gx_pred) + jnp.abs(gy_true - gy_pred)) * valid_mask\n    return jnp.sum(diff) / (jnp.sum(valid_mask) + 1e-6)\n\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# 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    # --- 1. FORCE CLEANUP & RE-PREPROCESSING ---\n    # We must clear the temp dir because previous runs left corrupt/mismatched files\n    if OUT_DIR.exists():\n        print(f\"🧹 Clearing corrupt cache at {OUT_DIR}...\")\n        shutil.rmtree(OUT_DIR)\n    OUT_DIR.mkdir(parents=True, exist_ok=True)\n    \n    # Run Preprocessing on ALL files (Avoid slicing [:10] to prevent mismatches)\n    run_preprocessing()\n    \n    # --- 2. VERIFY DATASET ---\n    # Extract IDs correctly\n    ids = sorted({p.stem.replace(\"_img\", \"\") for p in OUT_DIR.glob(\"*_img.npy\")})\n    print(f\"found {len(ids)} potential volume IDs.\")\n    \n    # Create Dataset\n    ds = VesuviusSurfaceDataset(ids, OUT_DIR, SAMPLES_PER_ID, train=True)\n    print(f\"📊 Dataset Size: {len(ds)} patches.\")\n    \n    if len(ds) == 0:\n        raise RuntimeError(\"❌ DATASET IS EMPTY! Check image/label filename matching.\")\n\n    # --- 3. TRAINING SETUP ---\n    # Use drop_last=False to ensure we use every bit of data\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\"🚀 Training Steps per Epoch: {steps_per_epoch}\")\n\n    # Setup Keys\n    rng = jax.random.PRNGKey(42)\n    rng_g, rng_d = jax.random.split(rng)\n\n    # Initialize Models\n    netG, netD = Pix2PixGenerator(), TopologyDiscriminator()\n    dummy_g = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 2*Z_CONTEXT+1))\n    dummy_target = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1))\n\n    v_g = netG.init(rng_g, dummy_g)\n    v_d = netD.init(rng_d, dummy_g, dummy_target)\n\n    # 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    # --- 4. OPTIMIZERS WITH CLIPPING (Prevents NaN) ---\n    tx_g = optax.chain(\n        optax.clip_by_global_norm(1.0),\n        optax.adam(sched_G, b1=0.5, b2=0.999)\n    )\n    tx_d = optax.chain(\n        optax.clip_by_global_norm(1.0),\n        optax.adam(sched_D, b1=0.5, b2=0.999)\n    )\n\n    # Create States\n    if LOADWeights:\n        print(f\"Loading weights from: {path_to_weights}\")\n        state_g = JaxTrainState.create(apply_fn=netG.apply, params=loaded_params, batch_stats=loaded_stats, tx=tx_g)\n    else:\n        state_g = JaxTrainState.create(apply_fn=netG.apply, params=v_g['params'], batch_stats=v_g['batch_stats'], tx=tx_g)\n\n    state_d = JaxTrainState.create(apply_fn=netD.apply, params=v_d['params'], batch_stats=v_d['batch_stats'], tx=tx_d)\n\n    # --- 5. TRAINING LOOP ---\n    best_loss = float('inf')\n    start_time = time.time()\n    \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        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            if step % 50 == 0:\n                print(f\"Ep {ep} | Step {step} | G: {lg:.4f} | D: {ld:.4f}\")\n\n        # Check if epoch was empty\n        if len(epoch_g_loss) == 0:\n            print(f\"⚠️ Epoch {ep} had NO DATA. check loader.\")\n            continue\n\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} | Avg G: {avg_g:.4f} | Avg D: {avg_d:.4f} | Time: {elapsed:.2f}h\")\n\n        if ep % 10 == 0:\n            save_portable_npy(state_g, f\"pix2pix_ep_{ep}.npz\")\n            visualize_results(state_g, ds, ep)\n            \n        if avg_g < best_loss:\n            best_loss = float(avg_g)\n            save_portable_npy(state_g, \"pix2pix_best.npz\")\n# if __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#     rng_g, rng_d = jax.random.split(rng)  # <--- DEFINES rng_g AND rng_d\n\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#     tx_g = optax.chain(\n#     optax.clip_by_global_norm(1.0),\n#     optax.adam(learning_rate=sched_G, b1=0.5, b2=0.999)\n#     )\n    \n#     tx_d = optax.chain(\n#         optax.clip_by_global_norm(1.0),\n#         optax.adam(learning_rate=sched_D, b1=0.5, b2=0.999)\n#     )\n#     # 2. DEFINE OPTIMIZER WITH CLIPPING (New Code)\n#     # This prevents the gradient explosion that causes NaNs\n#     optimizer_g = optax.chain(\n#         optax.clip_by_global_norm(1.0),  # <--- Crucial Fix: Cap gradients at 1.0\n#         optax.adam(sched_G, b1=0.5, b2=0.999)\n#     )\n    \n#     optimizer_d = optax.chain(\n#         optax.clip_by_global_norm(1.0),  # <--- Crucial Fix\n#         optax.adam(sched_D, b1=0.5, b2=0.999)\n#     )\n#     dummy_g = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 2*Z_CONTEXT+1))\n\n#     # Initialize weights randomly first\n#     v_g = netG.init(rng_g, dummy_g)\n#     # v_g = netG.init(jax.random.split(rng)[0], dummy)\n#     v_d = netD.init(jax.random.split(rng)[1], dummy, jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1)))\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, \n#             params=loaded_params, \n#             batch_stats=loaded_stats, \n#             tx=tx_g  # <--- Use new optimizer\n#         )\n#     else:\n#         state_g = JaxTrainState.create(\n#             apply_fn=netG.apply, \n#             params=v_g['params'], \n#             batch_stats=v_g['batch_stats'], \n#             tx=tx_g # <--- Use new optimizer\n#         )\n#         # Dummy target (Label) is 1 channel\n#     dummy_target = jnp.ones((1, PATCH_SIZE, PATCH_SIZE, 1)) \n    \n#     # Initialize Discriminator weights\n#     v_d = netD.init(rng_d, dummy_g, dummy_target)\n#     state_d = JaxTrainState.create(\n#         apply_fn=netD.apply, \n#         params=v_d['params'], \n#         batch_stats=v_d['batch_stats'], \n#         tx=tx_d # <--- Use new optimizer\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-15T00:45:14.367705Z","iopub.execute_input":"2026-02-15T00:45:14.367996Z","execution_failed":"2026-02-15T00:47:13.853Z"}},"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-15T00:44:30.560356Z","iopub.status.idle":"2026-02-15T00:44:30.560799Z","shell.execute_reply.started":"2026-02-15T00:44:30.560444Z","shell.execute_reply":"2026-02-15T00:44:30.560453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}